diff --git a/src/layer/arm/gemm_arm.cpp b/src/layer/arm/gemm_arm.cpp index 38528293aa7..636848deadf 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,158 @@ 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, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + 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); + 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; + + const int nn_MK = nn_M * nn_K; + #pragma omp parallel for num_threads(nT) + 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; + + 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(local_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, 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); + } + + 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 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); + 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 local_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(local_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, 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); + 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, 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()); + + 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); + 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; + Mat AT_tile_k(max_kk, max_ii, (signed char*)AT_tile + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile_k(local_block_count, max_ii, (float*)AT_descales_tile + (size_t)(k / block_size) * mr, (size_t)4u); + + if (j == 0) + { + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile_k, AT_descales_tile_k, i, max_ii, k, max_kk, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile_k, AT_descales_tile_k, i, max_ii, k, max_kk, block_size, input_scale_ptr); + } + + 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, 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); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + } + } + } + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_arm::create_pipeline(const Option& opt) { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + int weight_bits; + int block_size; + bool has_input_scale; + int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits == 8) + return create_pipeline_wq_int8(opt); +#endif return 0; } @@ -4745,10 +4897,181 @@ 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; + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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); + if (ret != 0) + return ret; + if (BT_data_packed.empty() || BT_data_packed_descales.empty()) + return -100; + + 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; + } + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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) + { + if (output_N1M) + top_blob.create(M, 1, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + } + else + { + if (output_N1M) + top_blob.create(N, 1, M, (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_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 + int weight_bits; + int block_size; + bool has_input_scale; + int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits == 8) + 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 7caf73e3876..f5ab207dd46 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 07b389bb4e9..f8c29efcf2c 100644 --- a/src/layer/arm/gemm_arm_asimddp.cpp +++ b/src/layer/arm/gemm_arm_asimddp.cpp @@ -3,6 +3,7 @@ #include "cpu.h" #include "mat.h" +#include "layer.h" #include "arm_usability.h" namespace ncnn { @@ -14,6 +15,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, const Option& opt) +{ + return pack_B_wq_int8(B, B_scales, BT, BT_descales, N, K, block_size, opt); +} + +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 max_kk, int block_size, const float* 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_asimddp(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, k, max_kk, 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 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, max_kk, 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 596fef8a248..fe116cc08ca 100644 --- a/src/layer/arm/gemm_arm_i8mm.cpp +++ b/src/layer/arm/gemm_arm_i8mm.cpp @@ -3,6 +3,7 @@ #include "cpu.h" #include "mat.h" +#include "layer.h" #include "arm_usability.h" namespace ncnn { @@ -14,6 +15,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, const Option& opt) +{ + return pack_B_wq_int8(B, B_scales, BT, BT_descales, N, K, block_size, opt); +} + +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 max_kk, int block_size, const float* 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_i8mm(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, k, max_kk, 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 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, max_kk, 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_wq_int8.h b/src/layer/arm/gemm_wq_int8.h new file mode 100644 index 00000000000..a3f030cbac6 --- /dev/null +++ b/src/layer/arm/gemm_wq_int8.h @@ -0,0 +1,7480 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#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, const Option& opt); +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 max_kk, 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 max_kk, 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 max_kk, 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, const Option& opt); +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 max_kk, 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 max_kk, 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 max_kk, 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 max_kk, 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, max_kk, 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, max_kk, 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 = (max_kk + 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) + { + signed char* pp = outptr + ii * out_hstep; + const float* pA0 = (const float*)A + (i + ii) * A_hstep + k; + const float* ps = input_scale_ptr ? input_scale_ptr + k : 0; + float* pd = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - g * block_size, block_size); + float32x4_t _absmax0 = vdupq_n_f32(0.f); + float32x4_t _absmax1 = vdupq_n_f32(0.f); + float32x4_t _absmax2 = vdupq_n_f32(0.f); + float32x4_t _absmax3 = vdupq_n_f32(0.f); + float32x4_t _absmax4 = vdupq_n_f32(0.f); + float32x4_t _absmax5 = vdupq_n_f32(0.f); + float32x4_t _absmax6 = vdupq_n_f32(0.f); + float32x4_t _absmax7 = vdupq_n_f32(0.f); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + float32x4_t _v0 = vld1q_f32(pA0 + kk); + float32x4_t _v1 = vld1q_f32(pA0 + A_hstep + kk); + float32x4_t _v2 = vld1q_f32(pA0 + A_hstep * 2 + kk); + float32x4_t _v3 = vld1q_f32(pA0 + A_hstep * 3 + kk); + float32x4_t _v4 = vld1q_f32(pA0 + A_hstep * 4 + kk); + float32x4_t _v5 = vld1q_f32(pA0 + A_hstep * 5 + kk); + float32x4_t _v6 = vld1q_f32(pA0 + A_hstep * 6 + kk); + float32x4_t _v7 = vld1q_f32(pA0 + A_hstep * 7 + kk); + if (ps) + { + float32x4_t _s = vld1q_f32(ps + kk); + _v0 = vmulq_f32(_v0, _s); + _v1 = vmulq_f32(_v1, _s); + _v2 = vmulq_f32(_v2, _s); + _v3 = vmulq_f32(_v3, _s); + _v4 = vmulq_f32(_v4, _s); + _v5 = vmulq_f32(_v5, _s); + _v6 = vmulq_f32(_v6, _s); + _v7 = vmulq_f32(_v7, _s); + } + _absmax0 = vmaxq_f32(_absmax0, vabsq_f32(_v0)); + _absmax1 = vmaxq_f32(_absmax1, vabsq_f32(_v1)); + _absmax2 = vmaxq_f32(_absmax2, vabsq_f32(_v2)); + _absmax3 = vmaxq_f32(_absmax3, vabsq_f32(_v3)); + _absmax4 = vmaxq_f32(_absmax4, vabsq_f32(_v4)); + _absmax5 = vmaxq_f32(_absmax5, vabsq_f32(_v5)); + _absmax6 = vmaxq_f32(_absmax6, vabsq_f32(_v6)); + _absmax7 = vmaxq_f32(_absmax7, vabsq_f32(_v7)); + } + float absmax0 = vmaxvq_f32(_absmax0); + float absmax1 = vmaxvq_f32(_absmax1); + float absmax2 = vmaxvq_f32(_absmax2); + float absmax3 = vmaxvq_f32(_absmax3); + float absmax4 = vmaxvq_f32(_absmax4); + float absmax5 = vmaxvq_f32(_absmax5); + float absmax6 = vmaxvq_f32(_absmax6); + float absmax7 = vmaxvq_f32(_absmax7); + for (; kk < max_kk0; kk++) + { + float v0 = pA0[kk]; + float v1 = pA0[A_hstep + kk]; + float v2 = pA0[A_hstep * 2 + kk]; + float v3 = pA0[A_hstep * 3 + kk]; + float v4 = pA0[A_hstep * 4 + kk]; + float v5 = pA0[A_hstep * 5 + kk]; + float v6 = pA0[A_hstep * 6 + kk]; + float v7 = pA0[A_hstep * 7 + kk]; + if (ps) + { + const float s = ps[kk]; + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + v4 *= s; + v5 *= s; + v6 *= s; + v7 *= 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)); + absmax4 = std::max(absmax4, fabsf(v4)); + absmax5 = std::max(absmax5, fabsf(v5)); + absmax6 = std::max(absmax6, fabsf(v6)); + absmax7 = std::max(absmax7, fabsf(v7)); + } + + 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; + 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; + + kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + float32x4_t _s0; + float32x4_t _s1; + if (ps) + { + _s0 = vld1q_f32(ps + kk); + _s1 = vld1q_f32(ps + kk + 4); + } + + float32x4_t _v00 = vld1q_f32(pA0 + kk); + float32x4_t _v01 = vld1q_f32(pA0 + kk + 4); + float32x4_t _v10 = vld1q_f32(pA0 + A_hstep + kk); + float32x4_t _v11 = vld1q_f32(pA0 + A_hstep + kk + 4); + float32x4_t _v20 = vld1q_f32(pA0 + A_hstep * 2 + kk); + float32x4_t _v21 = vld1q_f32(pA0 + A_hstep * 2 + kk + 4); + float32x4_t _v30 = vld1q_f32(pA0 + A_hstep * 3 + kk); + float32x4_t _v31 = vld1q_f32(pA0 + A_hstep * 3 + kk + 4); + float32x4_t _v40 = vld1q_f32(pA0 + A_hstep * 4 + kk); + float32x4_t _v41 = vld1q_f32(pA0 + A_hstep * 4 + kk + 4); + float32x4_t _v50 = vld1q_f32(pA0 + A_hstep * 5 + kk); + float32x4_t _v51 = vld1q_f32(pA0 + A_hstep * 5 + kk + 4); + float32x4_t _v60 = vld1q_f32(pA0 + A_hstep * 6 + kk); + float32x4_t _v61 = vld1q_f32(pA0 + A_hstep * 6 + kk + 4); + float32x4_t _v70 = vld1q_f32(pA0 + A_hstep * 7 + kk); + float32x4_t _v71 = vld1q_f32(pA0 + A_hstep * 7 + kk + 4); + if (ps) + { + _v00 = vmulq_f32(_v00, _s0); + _v01 = vmulq_f32(_v01, _s1); + _v10 = vmulq_f32(_v10, _s0); + _v11 = vmulq_f32(_v11, _s1); + _v20 = vmulq_f32(_v20, _s0); + _v21 = vmulq_f32(_v21, _s1); + _v30 = vmulq_f32(_v30, _s0); + _v31 = vmulq_f32(_v31, _s1); + _v40 = vmulq_f32(_v40, _s0); + _v41 = vmulq_f32(_v41, _s1); + _v50 = vmulq_f32(_v50, _s0); + _v51 = vmulq_f32(_v51, _s1); + _v60 = vmulq_f32(_v60, _s0); + _v61 = vmulq_f32(_v61, _s1); + _v70 = vmulq_f32(_v70, _s0); + _v71 = vmulq_f32(_v71, _s1); + } + int8x8_t _q0 = float2int8(vmulq_n_f32(_v00, scale0), vmulq_n_f32(_v01, scale0)); + int8x8_t _q1 = float2int8(vmulq_n_f32(_v10, scale1), vmulq_n_f32(_v11, scale1)); + int8x8_t _q2 = float2int8(vmulq_n_f32(_v20, scale2), vmulq_n_f32(_v21, scale2)); + int8x8_t _q3 = float2int8(vmulq_n_f32(_v30, scale3), vmulq_n_f32(_v31, scale3)); + int8x8_t _q4 = float2int8(vmulq_n_f32(_v40, scale4), vmulq_n_f32(_v41, scale4)); + int8x8_t _q5 = float2int8(vmulq_n_f32(_v50, scale5), vmulq_n_f32(_v51, scale5)); + int8x8_t _q6 = float2int8(vmulq_n_f32(_v60, scale6), vmulq_n_f32(_v61, scale6)); + int8x8_t _q7 = float2int8(vmulq_n_f32(_v70, scale7), vmulq_n_f32(_v71, scale7)); + vst1q_s8(pp, vcombine_s8(_q0, _q1)); + vst1q_s8(pp + 16, vcombine_s8(_q2, _q3)); + vst1q_s8(pp + 32, vcombine_s8(_q4, _q5)); + vst1q_s8(pp + 48, vcombine_s8(_q6, _q7)); + pp += 64; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + float32x4_t _v0 = vld1q_f32(pA0 + kk); + float32x4_t _v1 = vld1q_f32(pA0 + A_hstep + kk); + float32x4_t _v2 = vld1q_f32(pA0 + A_hstep * 2 + kk); + float32x4_t _v3 = vld1q_f32(pA0 + A_hstep * 3 + kk); + float32x4_t _v4 = vld1q_f32(pA0 + A_hstep * 4 + kk); + float32x4_t _v5 = vld1q_f32(pA0 + A_hstep * 5 + kk); + float32x4_t _v6 = vld1q_f32(pA0 + A_hstep * 6 + kk); + float32x4_t _v7 = vld1q_f32(pA0 + A_hstep * 7 + kk); + if (ps) + { + float32x4_t _s = vld1q_f32(ps + kk); + _v0 = vmulq_f32(_v0, _s); + _v1 = vmulq_f32(_v1, _s); + _v2 = vmulq_f32(_v2, _s); + _v3 = vmulq_f32(_v3, _s); + _v4 = vmulq_f32(_v4, _s); + _v5 = vmulq_f32(_v5, _s); + _v6 = vmulq_f32(_v6, _s); + _v7 = vmulq_f32(_v7, _s); + } + int8x8_t _q01 = float2int8(vmulq_n_f32(_v0, scale0), vmulq_n_f32(_v1, scale1)); + int8x8_t _q23 = float2int8(vmulq_n_f32(_v2, scale2), vmulq_n_f32(_v3, scale3)); + int8x8_t _q45 = float2int8(vmulq_n_f32(_v4, scale4), vmulq_n_f32(_v5, scale5)); + int8x8_t _q67 = float2int8(vmulq_n_f32(_v6, scale6), vmulq_n_f32(_v7, scale7)); +#if __ARM_FEATURE_DOTPROD + vst1q_s8(pp, vcombine_s8(_q01, _q23)); + vst1q_s8(pp + 16, vcombine_s8(_q45, _q67)); +#else + int16x8x2_t _q04 = vuzpq_s16(vreinterpretq_s16_s8(vcombine_s8(_q01, _q23)), vreinterpretq_s16_s8(vcombine_s8(_q45, _q67))); + vst1q_s16((short*)pp, _q04.val[0]); + vst1q_s16((short*)pp + 8, _q04.val[1]); +#endif + pp += 32; + } + for (; kk + 1 < max_kk0; kk += 2) + { + float s0 = 0.f; + float s1 = 0.f; + if (ps) + { + s0 = ps[kk]; + s1 = ps[kk + 1]; + } + float v00 = pA0[kk]; + float v01 = pA0[kk + 1]; + float v10 = pA0[A_hstep + kk]; + float v11 = pA0[A_hstep + kk + 1]; + float v20 = pA0[A_hstep * 2 + kk]; + float v21 = pA0[A_hstep * 2 + kk + 1]; + float v30 = pA0[A_hstep * 3 + kk]; + float v31 = pA0[A_hstep * 3 + kk + 1]; + float v40 = pA0[A_hstep * 4 + kk]; + float v41 = pA0[A_hstep * 4 + kk + 1]; + float v50 = pA0[A_hstep * 5 + kk]; + float v51 = pA0[A_hstep * 5 + kk + 1]; + float v60 = pA0[A_hstep * 6 + kk]; + float v61 = pA0[A_hstep * 6 + kk + 1]; + float v70 = pA0[A_hstep * 7 + kk]; + float v71 = pA0[A_hstep * 7 + kk + 1]; + if (ps) + { + v00 *= s0; + v01 *= s1; + v10 *= s0; + v11 *= s1; + v20 *= s0; + v21 *= s1; + v30 *= s0; + v31 *= s1; + v40 *= s0; + v41 *= s1; + v50 *= s0; + v51 *= s1; + v60 *= s0; + v61 *= s1; + v70 *= s0; + v71 *= s1; + } + *pp++ = float2int8(v00 * scale0); + *pp++ = float2int8(v01 * scale0); + *pp++ = float2int8(v10 * scale1); + *pp++ = float2int8(v11 * scale1); + *pp++ = float2int8(v20 * scale2); + *pp++ = float2int8(v21 * scale2); + *pp++ = float2int8(v30 * scale3); + *pp++ = float2int8(v31 * scale3); + *pp++ = float2int8(v40 * scale4); + *pp++ = float2int8(v41 * scale4); + *pp++ = float2int8(v50 * scale5); + *pp++ = float2int8(v51 * scale5); + *pp++ = float2int8(v60 * scale6); + *pp++ = float2int8(v61 * scale6); + *pp++ = float2int8(v70 * scale7); + *pp++ = float2int8(v71 * scale7); + } + if (kk < max_kk0) + { + float s = 0.f; + if (ps) + s = ps[kk]; + float v0 = pA0[kk]; + float v1 = pA0[A_hstep + kk]; + float v2 = pA0[A_hstep * 2 + kk]; + float v3 = pA0[A_hstep * 3 + kk]; + float v4 = pA0[A_hstep * 4 + kk]; + float v5 = pA0[A_hstep * 5 + kk]; + float v6 = pA0[A_hstep * 6 + kk]; + float v7 = pA0[A_hstep * 7 + kk]; + if (ps) + { + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + v4 *= s; + v5 *= s; + v6 *= s; + v7 *= s; + } + *pp++ = float2int8(v0 * scale0); + *pp++ = float2int8(v1 * scale1); + *pp++ = float2int8(v2 * scale2); + *pp++ = float2int8(v3 * scale3); + *pp++ = float2int8(v4 * scale4); + *pp++ = float2int8(v5 * scale5); + *pp++ = float2int8(v6 * scale6); + *pp++ = float2int8(v7 * scale7); + } + pA0 += max_kk0; + if (ps) + ps += max_kk0; + pd += 8; + } + } +#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 + k; + const float* ps = input_scale_ptr ? input_scale_ptr + k : 0; + + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - g * block_size, block_size); + float32x4_t _absmax0 = vdupq_n_f32(0.f); + float32x4_t _absmax1 = vdupq_n_f32(0.f); + float32x4_t _absmax2 = vdupq_n_f32(0.f); + float32x4_t _absmax3 = vdupq_n_f32(0.f); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + float32x4_t _v0 = vld1q_f32(pA0 + kk); + float32x4_t _v1 = vld1q_f32(pA0 + A_hstep + kk); + float32x4_t _v2 = vld1q_f32(pA0 + A_hstep * 2 + kk); + float32x4_t _v3 = vld1q_f32(pA0 + A_hstep * 3 + kk); + if (ps) + { + float32x4_t _s = vld1q_f32(ps + kk); + _v0 = vmulq_f32(_v0, _s); + _v1 = vmulq_f32(_v1, _s); + _v2 = vmulq_f32(_v2, _s); + _v3 = vmulq_f32(_v3, _s); + } + _absmax0 = vmaxq_f32(_absmax0, vabsq_f32(_v0)); + _absmax1 = vmaxq_f32(_absmax1, vabsq_f32(_v1)); + _absmax2 = vmaxq_f32(_absmax2, vabsq_f32(_v2)); + _absmax3 = vmaxq_f32(_absmax3, vabsq_f32(_v3)); + } +#if __aarch64__ + float absmax0 = vmaxvq_f32(_absmax0); + float absmax1 = vmaxvq_f32(_absmax1); + float absmax2 = vmaxvq_f32(_absmax2); + float absmax3 = vmaxvq_f32(_absmax3); +#else + float32x2_t _max0 = vmax_f32(vget_low_f32(_absmax0), vget_high_f32(_absmax0)); + float32x2_t _max1 = vmax_f32(vget_low_f32(_absmax1), vget_high_f32(_absmax1)); + float32x2_t _max2 = vmax_f32(vget_low_f32(_absmax2), vget_high_f32(_absmax2)); + float32x2_t _max3 = vmax_f32(vget_low_f32(_absmax3), vget_high_f32(_absmax3)); + _max0 = vpmax_f32(_max0, _max0); + _max1 = vpmax_f32(_max1, _max1); + _max2 = vpmax_f32(_max2, _max2); + _max3 = vpmax_f32(_max3, _max3); + float absmax0 = vget_lane_f32(_max0, 0); + float absmax1 = vget_lane_f32(_max1, 0); + float absmax2 = vget_lane_f32(_max2, 0); + float absmax3 = vget_lane_f32(_max3, 0); +#endif + for (; kk < max_kk0; kk++) + { + float v0 = pA0[kk]; + float v1 = pA0[A_hstep + kk]; + float v2 = pA0[A_hstep * 2 + kk]; + float v3 = pA0[A_hstep * 3 + kk]; + if (ps) + { + const float s = ps[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)); + } + descale_ptr[0] = absmax0 / 127.f; + descale_ptr[1] = absmax1 / 127.f; + descale_ptr[2] = absmax2 / 127.f; + descale_ptr[3] = absmax3 / 127.f; + 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; + + kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + float32x4_t _s0; + float32x4_t _s1; + if (ps) + { + _s0 = vld1q_f32(ps + kk); + _s1 = vld1q_f32(ps + kk + 4); + } + + float32x4_t _v00 = vld1q_f32(pA0 + kk); + float32x4_t _v01 = vld1q_f32(pA0 + kk + 4); + float32x4_t _v10 = vld1q_f32(pA0 + A_hstep + kk); + float32x4_t _v11 = vld1q_f32(pA0 + A_hstep + kk + 4); + float32x4_t _v20 = vld1q_f32(pA0 + A_hstep * 2 + kk); + float32x4_t _v21 = vld1q_f32(pA0 + A_hstep * 2 + kk + 4); + float32x4_t _v30 = vld1q_f32(pA0 + A_hstep * 3 + kk); + float32x4_t _v31 = vld1q_f32(pA0 + A_hstep * 3 + kk + 4); + if (ps) + { + _v00 = vmulq_f32(_v00, _s0); + _v01 = vmulq_f32(_v01, _s1); + _v10 = vmulq_f32(_v10, _s0); + _v11 = vmulq_f32(_v11, _s1); + _v20 = vmulq_f32(_v20, _s0); + _v21 = vmulq_f32(_v21, _s1); + _v30 = vmulq_f32(_v30, _s0); + _v31 = vmulq_f32(_v31, _s1); + } + int8x8_t _q0 = float2int8(vmulq_n_f32(_v00, scale0), vmulq_n_f32(_v01, scale0)); + int8x8_t _q1 = float2int8(vmulq_n_f32(_v10, scale1), vmulq_n_f32(_v11, scale1)); + int8x8_t _q2 = float2int8(vmulq_n_f32(_v20, scale2), vmulq_n_f32(_v21, scale2)); + int8x8_t _q3 = float2int8(vmulq_n_f32(_v30, scale3), vmulq_n_f32(_v31, scale3)); + vst1q_s8(pp, vcombine_s8(_q0, _q1)); + vst1q_s8(pp + 16, vcombine_s8(_q2, _q3)); + pp += 32; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + float32x4_t _v0 = vld1q_f32(pA0 + kk); + float32x4_t _v1 = vld1q_f32(pA0 + A_hstep + kk); + float32x4_t _v2 = vld1q_f32(pA0 + A_hstep * 2 + kk); + float32x4_t _v3 = vld1q_f32(pA0 + A_hstep * 3 + kk); + if (ps) + { + float32x4_t _s = vld1q_f32(ps + kk); + _v0 = vmulq_f32(_v0, _s); + _v1 = vmulq_f32(_v1, _s); + _v2 = vmulq_f32(_v2, _s); + _v3 = vmulq_f32(_v3, _s); + } + int8x8_t _q01 = float2int8(vmulq_n_f32(_v0, scale0), vmulq_n_f32(_v1, scale1)); + int8x8_t _q23 = float2int8(vmulq_n_f32(_v2, scale2), vmulq_n_f32(_v3, scale3)); +#if __ARM_FEATURE_DOTPROD + vst1q_s8(pp, vcombine_s8(_q01, _q23)); +#else + int16x8_t _q0123 = vreinterpretq_s16_s8(vcombine_s8(_q01, _q23)); + int16x8x2_t _q02 = vuzpq_s16(_q0123, _q0123); + vst1q_s8(pp, vreinterpretq_s8_s16(vcombine_s16(vget_low_s16(_q02.val[0]), vget_low_s16(_q02.val[1])))); +#endif + pp += 16; + } + for (; kk + 1 < max_kk0; kk += 2) + { + float s0 = 0.f; + float s1 = 0.f; + if (ps) + { + s0 = ps[kk]; + s1 = ps[kk + 1]; + } + float v00 = pA0[kk]; + float v01 = pA0[kk + 1]; + float v10 = pA0[A_hstep + kk]; + float v11 = pA0[A_hstep + kk + 1]; + float v20 = pA0[A_hstep * 2 + kk]; + float v21 = pA0[A_hstep * 2 + kk + 1]; + float v30 = pA0[A_hstep * 3 + kk]; + float v31 = pA0[A_hstep * 3 + kk + 1]; + if (ps) + { + v00 *= s0; + v01 *= s1; + v10 *= s0; + v11 *= s1; + v20 *= s0; + v21 *= s1; + v30 *= s0; + v31 *= s1; + } + *pp++ = float2int8(v00 * scale0); + *pp++ = float2int8(v01 * scale0); + *pp++ = float2int8(v10 * scale1); + *pp++ = float2int8(v11 * scale1); + *pp++ = float2int8(v20 * scale2); + *pp++ = float2int8(v21 * scale2); + *pp++ = float2int8(v30 * scale3); + *pp++ = float2int8(v31 * scale3); + } + if (kk < max_kk0) + { + float s = 0.f; + if (ps) + s = ps[kk]; + float v0 = pA0[kk]; + float v1 = pA0[A_hstep + kk]; + float v2 = pA0[A_hstep * 2 + kk]; + float v3 = pA0[A_hstep * 3 + kk]; + if (ps) + { + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + } + *pp++ = float2int8(v0 * scale0); + *pp++ = float2int8(v1 * scale1); + *pp++ = float2int8(v2 * scale2); + *pp++ = float2int8(v3 * scale3); + } + pA0 += max_kk0; + if (ps) + ps += max_kk0; + descale_ptr += 4; + } + } +#endif // __ARM_NEON + for (; ii + 1 < max_ii; ii += 2) + { + signed char* pp = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + const float* pA0g = (const float*)A + (i + ii) * A_hstep + k; + const float* pA1g = pA0g + A_hstep; + const float* ps = input_scale_ptr ? input_scale_ptr + k : 0; + + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - g * block_size, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + int kk = 0; +#if __ARM_NEON + float32x4_t _absmax0 = vdupq_n_f32(0.f); + float32x4_t _absmax1 = vdupq_n_f32(0.f); + for (; kk + 3 < max_kk0; kk += 4) + { + float32x4_t _v0 = vld1q_f32(pA0g + kk); + float32x4_t _v1 = vld1q_f32(pA1g + kk); + if (ps) + { + float32x4_t _s = vld1q_f32(ps + kk); + _v0 = vmulq_f32(_v0, _s); + _v1 = vmulq_f32(_v1, _s); + } + _absmax0 = vmaxq_f32(_absmax0, vabsq_f32(_v0)); + _absmax1 = vmaxq_f32(_absmax1, vabsq_f32(_v1)); + } +#if __aarch64__ + absmax0 = vmaxvq_f32(_absmax0); + absmax1 = vmaxvq_f32(_absmax1); +#else + float32x2_t _max0 = vmax_f32(vget_low_f32(_absmax0), vget_high_f32(_absmax0)); + float32x2_t _max1 = vmax_f32(vget_low_f32(_absmax1), vget_high_f32(_absmax1)); + _max0 = vpmax_f32(_max0, _max0); + _max1 = vpmax_f32(_max1, _max1); + absmax0 = vget_lane_f32(_max0, 0); + absmax1 = vget_lane_f32(_max1, 0); +#endif +#endif // __ARM_NEON + + for (; kk < max_kk0; kk++) + { + float v0 = pA0g[kk]; + float v1 = pA1g[kk]; + if (ps) + { + const float s = ps[kk]; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } + + 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; + + kk = 0; +#if __ARM_NEON +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + float32x4_t _v00 = vld1q_f32(pA0g + kk); + float32x4_t _v01 = vld1q_f32(pA0g + kk + 4); + float32x4_t _v10 = vld1q_f32(pA1g + kk); + float32x4_t _v11 = vld1q_f32(pA1g + kk + 4); + if (ps) + { + 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); + } + int8x8_t _q0 = float2int8(vmulq_n_f32(_v00, scale0), vmulq_n_f32(_v01, scale0)); + int8x8_t _q1 = float2int8(vmulq_n_f32(_v10, scale1), vmulq_n_f32(_v11, scale1)); + vst1q_s8(pp, vcombine_s8(_q0, _q1)); + pp += 16; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + float32x4_t _v0 = vld1q_f32(pA0g + kk); + float32x4_t _v1 = vld1q_f32(pA1g + kk); + if (ps) + { + float32x4_t _s = vld1q_f32(ps + kk); + _v0 = vmulq_f32(_v0, _s); + _v1 = vmulq_f32(_v1, _s); + } + int8x8_t _q01 = float2int8(vmulq_n_f32(_v0, scale0), vmulq_n_f32(_v1, scale1)); +#if __ARM_FEATURE_DOTPROD + vst1_s8(pp, _q01); +#else + int16x4_t _q01_s16 = vreinterpret_s16_s8(_q01); + int16x4_t _q10_s16 = vext_s16(_q01_s16, _q01_s16, 2); + vst1_s8(pp, vreinterpret_s8_s16(vzip_s16(_q01_s16, _q10_s16).val[0])); +#endif + pp += 8; + } + for (; kk + 1 < max_kk0; kk += 2) + { + float v00 = pA0g[kk]; + float v01 = pA0g[kk + 1]; + float v10 = pA1g[kk]; + float v11 = pA1g[kk + 1]; + if (ps) + { + v00 *= ps[kk]; + v01 *= ps[kk + 1]; + v10 *= ps[kk]; + v11 *= ps[kk + 1]; + } + *pp++ = float2int8(v00 * scale0); + *pp++ = float2int8(v01 * scale0); + *pp++ = float2int8(v10 * scale1); + *pp++ = float2int8(v11 * scale1); + } +#endif // __ARM_NEON + for (; kk < max_kk0; kk++) + { + float v0 = pA0g[kk]; + float v1 = pA1g[kk]; + if (ps) + { + v0 *= ps[kk]; + v1 *= ps[kk]; + } + *pp++ = float2int8(v0 * scale0); + *pp++ = float2int8(v1 * scale1); + } + + pA0g += max_kk0; + pA1g += max_kk0; + if (ps) + ps += max_kk0; + descale_ptr += 2; + } + } + for (; ii < max_ii; ii++) + { + const float* pAg = (const float*)A + (i + ii) * A_hstep + k; + const float* ps = input_scale_ptr ? input_scale_ptr + k : 0; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - g * block_size, block_size); + const float* pA = pAg; + const float* pscale = ps; + + float absmax = 0.f; + int kk = 0; +#if __ARM_NEON + float32x4_t _absmax = vdupq_n_f32(0.f); + for (; kk + 3 < max_kk0; kk += 4) + { + 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__ + 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_kk0; kk++) + { + float v = *pA++; + if (pscale) + v *= *pscale++; + absmax = std::max(absmax, fabsf(v)); + } + + if (absmax == 0.f) + { + *pd++ = 0.f; + for (int kk0 = 0; kk0 < max_kk0; kk0++) + *pp++ = 0; + pAg += max_kk0; + if (ps) + ps += max_kk0; + continue; + } + + const float scale = 127.f / absmax; + *pd++ = absmax / 127.f; + + kk = 0; + pA = pAg; + pscale = ps; +#if __ARM_NEON + float32x4_t _scale = vdupq_n_f32(scale); + for (; kk + 7 < max_kk0; kk += 8) + { + float32x4_t _v0 = vld1q_f32(pA); + float32x4_t _v1 = vld1q_f32(pA + 4); + pA += 8; + if (pscale) + { + _v0 = vmulq_f32(_v0, vld1q_f32(pscale)); + _v1 = vmulq_f32(_v1, vld1q_f32(pscale + 4)); + pscale += 8; + } + vst1_s8(pp, float2int8(vmulq_f32(_v0, _scale), vmulq_f32(_v1, _scale))); + pp += 8; + } + for (; kk + 3 < max_kk0; kk += 4) + { + float32x4_t _v = vld1q_f32(pA); + pA += 4; + if (pscale) + { + _v = vmulq_f32(_v, vld1q_f32(pscale)); + pscale += 4; + } + 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_kk0; kk++) + { + float v = *pA++; + if (pscale) + v *= *pscale++; + *pp++ = float2int8(v * scale); + } + + pAg += max_kk0; + if (ps) + ps += max_kk0; + } + } +} + +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_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, max_kk, 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, max_kk, 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 = (max_kk + 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 + (size_t)k * A_hstep + i + ii; + signed char* pp = outptr + ii * out_hstep; + const float* ptrAg = ptrA; + const float* ps = input_scale_ptr ? input_scale_ptr + k : 0; + float* pd = descales + ii * descales_hstep; + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - 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_kk0; kk++) + { + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + 4); + if (ps) + { + _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 absmax0 = vgetq_lane_f32(_absmax0, 0); + float absmax1 = vgetq_lane_f32(_absmax0, 1); + float absmax2 = vgetq_lane_f32(_absmax0, 2); + float absmax3 = vgetq_lane_f32(_absmax0, 3); + float absmax4 = vgetq_lane_f32(_absmax1, 0); + float absmax5 = vgetq_lane_f32(_absmax1, 1); + float absmax6 = vgetq_lane_f32(_absmax1, 2); + float absmax7 = vgetq_lane_f32(_absmax1, 3); + 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; + + float32x4_t _scale0 = vdupq_n_f32(absmax0 == 0.f ? 0.f : 127.f / absmax0); + _scale0 = vsetq_lane_f32(absmax1 == 0.f ? 0.f : 127.f / absmax1, _scale0, 1); + _scale0 = vsetq_lane_f32(absmax2 == 0.f ? 0.f : 127.f / absmax2, _scale0, 2); + _scale0 = vsetq_lane_f32(absmax3 == 0.f ? 0.f : 127.f / absmax3, _scale0, 3); + float32x4_t _scale1 = vdupq_n_f32(absmax4 == 0.f ? 0.f : 127.f / absmax4); + _scale1 = vsetq_lane_f32(absmax5 == 0.f ? 0.f : 127.f / absmax5, _scale1, 1); + _scale1 = vsetq_lane_f32(absmax6 == 0.f ? 0.f : 127.f / absmax6, _scale1, 2); + _scale1 = vsetq_lane_f32(absmax7 == 0.f ? 0.f : 127.f / absmax7, _scale1, 3); + + int kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v00 = vld1q_f32(ptrAk); + float32x4_t _v01 = vld1q_f32(ptrAk + 4); + float32x4_t _v10 = vld1q_f32(ptrAk + A_hstep); + float32x4_t _v11 = vld1q_f32(ptrAk + A_hstep + 4); + float32x4_t _v20 = vld1q_f32(ptrAk + A_hstep * 2); + float32x4_t _v21 = vld1q_f32(ptrAk + A_hstep * 2 + 4); + float32x4_t _v30 = vld1q_f32(ptrAk + A_hstep * 3); + float32x4_t _v31 = vld1q_f32(ptrAk + A_hstep * 3 + 4); + float32x4_t _v40 = vld1q_f32(ptrAk + A_hstep * 4); + float32x4_t _v41 = vld1q_f32(ptrAk + A_hstep * 4 + 4); + float32x4_t _v50 = vld1q_f32(ptrAk + A_hstep * 5); + float32x4_t _v51 = vld1q_f32(ptrAk + A_hstep * 5 + 4); + float32x4_t _v60 = vld1q_f32(ptrAk + A_hstep * 6); + float32x4_t _v61 = vld1q_f32(ptrAk + A_hstep * 6 + 4); + float32x4_t _v70 = vld1q_f32(ptrAk + A_hstep * 7); + float32x4_t _v71 = vld1q_f32(ptrAk + A_hstep * 7 + 4); + if (ps) + { + _v00 = vmulq_n_f32(_v00, ps[kk]); + _v01 = vmulq_n_f32(_v01, ps[kk]); + _v10 = vmulq_n_f32(_v10, ps[kk + 1]); + _v11 = vmulq_n_f32(_v11, ps[kk + 1]); + _v20 = vmulq_n_f32(_v20, ps[kk + 2]); + _v21 = vmulq_n_f32(_v21, ps[kk + 2]); + _v30 = vmulq_n_f32(_v30, ps[kk + 3]); + _v31 = vmulq_n_f32(_v31, ps[kk + 3]); + _v40 = vmulq_n_f32(_v40, ps[kk + 4]); + _v41 = vmulq_n_f32(_v41, ps[kk + 4]); + _v50 = vmulq_n_f32(_v50, ps[kk + 5]); + _v51 = vmulq_n_f32(_v51, ps[kk + 5]); + _v60 = vmulq_n_f32(_v60, ps[kk + 6]); + _v61 = vmulq_n_f32(_v61, ps[kk + 6]); + _v70 = vmulq_n_f32(_v70, ps[kk + 7]); + _v71 = vmulq_n_f32(_v71, ps[kk + 7]); + } + int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale1)); + int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale0), vmulq_f32(_v11, _scale1)); + int8x8_t _q2 = float2int8(vmulq_f32(_v20, _scale0), vmulq_f32(_v21, _scale1)); + int8x8_t _q3 = float2int8(vmulq_f32(_v30, _scale0), vmulq_f32(_v31, _scale1)); + int8x8_t _q4 = float2int8(vmulq_f32(_v40, _scale0), vmulq_f32(_v41, _scale1)); + int8x8_t _q5 = float2int8(vmulq_f32(_v50, _scale0), vmulq_f32(_v51, _scale1)); + int8x8_t _q6 = float2int8(vmulq_f32(_v60, _scale0), vmulq_f32(_v61, _scale1)); + int8x8_t _q7 = float2int8(vmulq_f32(_v70, _scale0), vmulq_f32(_v71, _scale1)); + int8x8x2_t _r04 = vzip_s8(_q0, _q4); + int8x8x2_t _r15 = vzip_s8(_q1, _q5); + int8x8x2_t _r26 = vzip_s8(_q2, _q6); + int8x8x2_t _r37 = vzip_s8(_q3, _q7); + 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 +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v00 = vld1q_f32(ptrAk); + float32x4_t _v01 = vld1q_f32(ptrAk + 4); + float32x4_t _v10 = vld1q_f32(ptrAk + A_hstep); + float32x4_t _v11 = vld1q_f32(ptrAk + A_hstep + 4); + float32x4_t _v20 = vld1q_f32(ptrAk + A_hstep * 2); + float32x4_t _v21 = vld1q_f32(ptrAk + A_hstep * 2 + 4); + float32x4_t _v30 = vld1q_f32(ptrAk + A_hstep * 3); + float32x4_t _v31 = vld1q_f32(ptrAk + A_hstep * 3 + 4); + if (ps) + { + _v00 = vmulq_n_f32(_v00, ps[kk]); + _v01 = vmulq_n_f32(_v01, ps[kk]); + _v10 = vmulq_n_f32(_v10, ps[kk + 1]); + _v11 = vmulq_n_f32(_v11, ps[kk + 1]); + _v20 = vmulq_n_f32(_v20, ps[kk + 2]); + _v21 = vmulq_n_f32(_v21, ps[kk + 2]); + _v30 = vmulq_n_f32(_v30, ps[kk + 3]); + _v31 = vmulq_n_f32(_v31, ps[kk + 3]); + } + int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale1)); + int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale0), vmulq_f32(_v11, _scale1)); + int8x8_t _q2 = float2int8(vmulq_f32(_v20, _scale0), vmulq_f32(_v21, _scale1)); + int8x8_t _q3 = float2int8(vmulq_f32(_v30, _scale0), vmulq_f32(_v31, _scale1)); +#if __ARM_FEATURE_DOTPROD + int8x8x4_t _r0123; + _r0123.val[0] = _q0; + _r0123.val[1] = _q1; + _r0123.val[2] = _q2; + _r0123.val[3] = _q3; + vst4_s8(pp, _r0123); +#else + int8x8x2_t _r01; + _r01.val[0] = _q0; + _r01.val[1] = _q1; + int8x8x2_t _r23; + _r23.val[0] = _q2; + _r23.val[1] = _q3; + vst2_s8(pp, _r01); + vst2_s8(pp + 16, _r23); +#endif + pp += 32; + } + for (; kk + 1 < max_kk0; kk += 2) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v00 = vld1q_f32(ptrAk); + float32x4_t _v01 = vld1q_f32(ptrAk + 4); + float32x4_t _v10 = vld1q_f32(ptrAk + A_hstep); + float32x4_t _v11 = vld1q_f32(ptrAk + A_hstep + 4); + if (ps) + { + _v00 = vmulq_n_f32(_v00, ps[kk]); + _v01 = vmulq_n_f32(_v01, ps[kk]); + _v10 = vmulq_n_f32(_v10, ps[kk + 1]); + _v11 = vmulq_n_f32(_v11, ps[kk + 1]); + } + int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale1)); + int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale0), vmulq_f32(_v11, _scale1)); + int8x8x2_t _r01; + _r01.val[0] = _q0; + _r01.val[1] = _q1; + vst2_s8(pp, _r01); + pp += 16; + } + if (kk < max_kk0) + { + 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, 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_kk0 * A_hstep; + if (ps) + ps += max_kk0; + pd += 8; + } + } +#endif // __aarch64__ + for (; ii + 3 < max_ii; ii += 4) + { + const float* ptrA = (const float*)A + (size_t)k * A_hstep + 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 ? input_scale_ptr + k : 0; + + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - g * block_size, block_size); + + float32x4_t _absmax = vdupq_n_f32(0.f); + const float* ptrAk = ptrAg; + for (int kk = 0; kk < max_kk0; kk++) + { + float32x4_t _v = vld1q_f32(ptrAk); + if (ps) + _v = vmulq_n_f32(_v, ps[kk]); + _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); + ptrAk += A_hstep; + } + + vst1q_f32(descale_ptr, vmulq_n_f32(_absmax, 1.f / 127.f)); + float absmax0 = vgetq_lane_f32(_absmax, 0); + float absmax1 = vgetq_lane_f32(_absmax, 1); + float absmax2 = vgetq_lane_f32(_absmax, 2); + float absmax3 = vgetq_lane_f32(_absmax, 3); + float32x4_t _scale = vdupq_n_f32(absmax0 == 0.f ? 0.f : 127.f / absmax0); + _scale = vsetq_lane_f32(absmax1 == 0.f ? 0.f : 127.f / absmax1, _scale, 1); + _scale = vsetq_lane_f32(absmax2 == 0.f ? 0.f : 127.f / absmax2, _scale, 2); + _scale = vsetq_lane_f32(absmax3 == 0.f ? 0.f : 127.f / absmax3, _scale, 3); + + int kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + A_hstep); + float32x4_t _v2 = vld1q_f32(ptrAk + A_hstep * 2); + float32x4_t _v3 = vld1q_f32(ptrAk + A_hstep * 3); + float32x4_t _v4 = vld1q_f32(ptrAk + A_hstep * 4); + float32x4_t _v5 = vld1q_f32(ptrAk + A_hstep * 5); + float32x4_t _v6 = vld1q_f32(ptrAk + A_hstep * 6); + float32x4_t _v7 = vld1q_f32(ptrAk + A_hstep * 7); + if (ps) + { + _v0 = vmulq_n_f32(_v0, ps[kk]); + _v1 = vmulq_n_f32(_v1, ps[kk + 1]); + _v2 = vmulq_n_f32(_v2, ps[kk + 2]); + _v3 = vmulq_n_f32(_v3, ps[kk + 3]); + _v4 = vmulq_n_f32(_v4, ps[kk + 4]); + _v5 = vmulq_n_f32(_v5, ps[kk + 5]); + _v6 = vmulq_n_f32(_v6, ps[kk + 6]); + _v7 = vmulq_n_f32(_v7, ps[kk + 7]); + } + _v0 = vmulq_f32(_v0, _scale); + _v1 = vmulq_f32(_v1, _scale); + _v2 = vmulq_f32(_v2, _scale); + _v3 = vmulq_f32(_v3, _scale); + _v4 = vmulq_f32(_v4, _scale); + _v5 = vmulq_f32(_v5, _scale); + _v6 = vmulq_f32(_v6, _scale); + _v7 = vmulq_f32(_v7, _scale); + float32x4x2_t _v04 = vzipq_f32(_v0, _v4); + float32x4x2_t _v15 = vzipq_f32(_v1, _v5); + float32x4x2_t _v26 = vzipq_f32(_v2, _v6); + float32x4x2_t _v37 = vzipq_f32(_v3, _v7); + int8x8x4_t _r0123; + _r0123.val[0] = float2int8(_v04.val[0], _v04.val[1]); + _r0123.val[1] = float2int8(_v15.val[0], _v15.val[1]); + _r0123.val[2] = float2int8(_v26.val[0], _v26.val[1]); + _r0123.val[3] = float2int8(_v37.val[0], _v37.val[1]); + vst4_s8(pp, _r0123); + pp += 32; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + A_hstep); + float32x4_t _v2 = vld1q_f32(ptrAk + A_hstep * 2); + float32x4_t _v3 = vld1q_f32(ptrAk + A_hstep * 3); + if (ps) + { + _v0 = vmulq_n_f32(_v0, ps[kk]); + _v1 = vmulq_n_f32(_v1, ps[kk + 1]); + _v2 = vmulq_n_f32(_v2, ps[kk + 2]); + _v3 = vmulq_n_f32(_v3, ps[kk + 3]); + } + _v0 = vmulq_f32(_v0, _scale); + _v1 = vmulq_f32(_v1, _scale); + _v2 = vmulq_f32(_v2, _scale); + _v3 = vmulq_f32(_v3, _scale); +#if __ARM_FEATURE_DOTPROD + transpose4x4_ps(_v0, _v1, _v2, _v3); + int8x8_t _q01 = float2int8(_v0, _v1); + int8x8_t _q23 = float2int8(_v2, _v3); + vst1q_s8(pp, vcombine_s8(_q01, _q23)); +#else + int8x8_t _q01 = float2int8(_v0, _v1); + int8x8_t _q23 = float2int8(_v2, _v3); + int8x8_t _q10 = vext_s8(_q01, _q01, 4); + int8x8_t _q32 = vext_s8(_q23, _q23, 4); + vst1_s8(pp, vzip_s8(_q01, _q10).val[0]); + vst1_s8(pp + 8, vzip_s8(_q23, _q32).val[0]); +#endif + pp += 16; + } + for (; kk + 1 < max_kk0; kk += 2) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + A_hstep); + if (ps) + { + _v0 = vmulq_n_f32(_v0, ps[kk]); + _v1 = vmulq_n_f32(_v1, ps[kk + 1]); + } + int8x8_t _q01 = float2int8(vmulq_f32(_v0, _scale), vmulq_f32(_v1, _scale)); + int8x8_t _q10 = vext_s8(_q01, _q01, 4); + vst1_s8(pp, vzip_s8(_q01, _q10).val[0]); + pp += 8; + } + if (kk < max_kk0) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v = vld1q_f32(ptrAk); + if (ps) + { + _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_kk0 * A_hstep; + if (ps) + ps += max_kk0; + descale_ptr += 4; + } + } +#endif // __ARM_NEON + for (; ii + 1 < max_ii; ii += 2) + { + const float* ptrA = (const float*)A + (size_t)k * A_hstep + 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 ? input_scale_ptr + k : 0; + + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - g * block_size, block_size); +#if __ARM_NEON + float32x2_t _absmax = vdup_n_f32(0.f); + const float* ptrAk = ptrAg; + for (int kk = 0; kk < max_kk0; kk++) + { + 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, vmul_n_f32(_absmax, 1.f / 127.f)); + float absmax0 = vget_lane_f32(_absmax, 0); + float absmax1 = vget_lane_f32(_absmax, 1); + float32x2_t _scale = vdup_n_f32(absmax0 == 0.f ? 0.f : 127.f / absmax0); + _scale = vset_lane_f32(absmax1 == 0.f ? 0.f : 127.f / absmax1, _scale, 1); + + int kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x2_t _v0 = vld1_f32(ptrAk); + float32x2_t _v1 = vld1_f32(ptrAk + A_hstep); + float32x2_t _v2 = vld1_f32(ptrAk + A_hstep * 2); + float32x2_t _v3 = vld1_f32(ptrAk + A_hstep * 3); + float32x2_t _v4 = vld1_f32(ptrAk + A_hstep * 4); + float32x2_t _v5 = vld1_f32(ptrAk + A_hstep * 5); + float32x2_t _v6 = vld1_f32(ptrAk + A_hstep * 6); + float32x2_t _v7 = vld1_f32(ptrAk + A_hstep * 7); + if (ps) + { + _v0 = vmul_n_f32(_v0, ps[kk]); + _v1 = vmul_n_f32(_v1, ps[kk + 1]); + _v2 = vmul_n_f32(_v2, ps[kk + 2]); + _v3 = vmul_n_f32(_v3, ps[kk + 3]); + _v4 = vmul_n_f32(_v4, ps[kk + 4]); + _v5 = vmul_n_f32(_v5, ps[kk + 5]); + _v6 = vmul_n_f32(_v6, ps[kk + 6]); + _v7 = vmul_n_f32(_v7, ps[kk + 7]); + } + float32x4_t _scale_scale = vcombine_f32(_scale, _scale); + float32x4_t _v01 = vmulq_f32(vcombine_f32(_v0, _v1), _scale_scale); + float32x4_t _v23 = vmulq_f32(vcombine_f32(_v2, _v3), _scale_scale); + float32x4_t _v45 = vmulq_f32(vcombine_f32(_v4, _v5), _scale_scale); + float32x4_t _v67 = vmulq_f32(vcombine_f32(_v6, _v7), _scale_scale); + int8x8_t _q0 = float2int8(_v01, _v23); + int8x8_t _q1 = float2int8(_v45, _v67); + int8x8x2_t _q01 = vuzp_s8(_q0, _q1); + vst1q_s8(pp, vcombine_s8(_q01.val[0], _q01.val[1])); + pp += 16; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x2_t _v0 = vld1_f32(ptrAk); + float32x2_t _v1 = vld1_f32(ptrAk + A_hstep); + float32x2_t _v2 = vld1_f32(ptrAk + A_hstep * 2); + float32x2_t _v3 = vld1_f32(ptrAk + A_hstep * 3); + if (ps) + { + _v0 = vmul_n_f32(_v0, ps[kk]); + _v1 = vmul_n_f32(_v1, ps[kk + 1]); + _v2 = vmul_n_f32(_v2, ps[kk + 2]); + _v3 = vmul_n_f32(_v3, ps[kk + 3]); + } + float32x4_t _scale_scale = vcombine_f32(_scale, _scale); + float32x4_t _v01 = vmulq_f32(vcombine_f32(_v0, _v1), _scale_scale); + float32x4_t _v23 = vmulq_f32(vcombine_f32(_v2, _v3), _scale_scale); + int8x8_t _q0123 = float2int8(_v01, _v23); +#if __ARM_FEATURE_DOTPROD + int8x8x2_t _q02 = vuzp_s8(_q0123, _q0123); + vst1_s8(pp, vreinterpret_s8_s32(vzip_s32(vreinterpret_s32_s8(_q02.val[0]), vreinterpret_s32_s8(_q02.val[1])).val[0])); +#else + int8x8x2_t _q02 = vuzp_s8(_q0123, _q0123); + vst1_s8(pp, vreinterpret_s8_s16(vzip_s16(vreinterpret_s16_s8(_q02.val[0]), vreinterpret_s16_s8(_q02.val[1])).val[0])); +#endif + pp += 8; + } + for (; kk + 1 < max_kk0; kk += 2) + { + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x2_t _v0 = vld1_f32(ptrAk); + float32x2_t _v1 = vld1_f32(ptrAk + A_hstep); + if (ps) + { + _v0 = vmul_n_f32(_v0, ps[kk]); + _v1 = vmul_n_f32(_v1, ps[kk + 1]); + } + float32x4_t _scale_scale = vcombine_f32(_scale, _scale); + float32x4_t _v01 = vmulq_f32(vcombine_f32(_v0, _v1), _scale_scale); + int8x8_t _q01 = float2int8(_v01, _v01); + int8x8x2_t _q02 = vuzp_s8(_q01, _q01); + vst1_lane_s16((short*)pp, vreinterpret_s16_s8(_q02.val[0]), 0); + vst1_lane_s16((short*)pp + 1, vreinterpret_s16_s8(_q02.val[1]), 0); + pp += 4; + } + if (kk < max_kk0) + { + 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; + } +#else + const float* ptrAk = ptrAg; + float absmax0 = 0.f; + float absmax1 = 0.f; + + for (int kk = 0; kk < max_kk0; kk++) + { + float v0 = ptrAk[0]; + float v1 = ptrAk[1]; + if (ps) + { + v0 *= ps[kk]; + v1 *= ps[kk]; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + ptrAk += A_hstep; + } + + 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 = ptrAg; + for (int kk = 0; kk < max_kk0; kk++) + { + float v0 = ptrAk[0]; + float v1 = ptrAk[1]; + if (ps) + { + v0 *= ps[kk]; + v1 *= ps[kk]; + } + *pp++ = float2int8(v0 * scale0); + *pp++ = float2int8(v1 * scale1); + ptrAk += A_hstep; + } +#endif // __ARM_NEON + ptrAg += (size_t)max_kk0 * A_hstep; + if (ps) + ps += max_kk0; + descale_ptr += 2; + } + } + for (; ii < max_ii; ii++) + { + const float* ptrA = (const float*)A + (size_t)k * A_hstep + 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 ? input_scale_ptr + k : 0; + + for (int g = 0; g < block_count; g++) + { + const int max_kk0 = std::min(max_kk - g * block_size, block_size); + + float absmax = 0.f; + const float* ptrAk = ptrAg; + for (int kk = 0; kk < max_kk0; kk++) + { + float v = *ptrAk; + if (ps) + v *= ps[kk]; + absmax = std::max(absmax, fabsf(v)); + ptrAk += A_hstep; + } + + if (absmax == 0.f) + { + *descale_ptr++ = 0.f; + for (int kk0 = 0; kk0 < max_kk0; kk0++) + *pp++ = 0; + ptrAg += (size_t)max_kk0 * A_hstep; + if (ps) + ps += max_kk0; + continue; + } + + const float scale = 127.f / absmax; + *descale_ptr++ = absmax / 127.f; + + ptrAk = ptrAg; + for (int kk = 0; kk < max_kk0; kk++) + { + float v = *ptrAk; + if (ps) + { + v *= ps[kk]; + } + *pp++ = float2int8(v * scale); + ptrAk += A_hstep; + } + ptrAg += (size_t)max_kk0 * A_hstep; + if (ps) + ps += max_kk0; + } + } +} + +// 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. +// The non-neon simd32 nr2 panel follows the ordinary gemm_int8 per-k producer. +// Its bytes are b0k0 b1k0 b0k1 b1k1 so sxtb16 and ror+sxtb16 extract the two +// output columns without an extra byte shuffle in the smlad inner loop. +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, const Option& opt) +{ +#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, opt); + } +#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, opt); + } +#endif + + const int block_count = (K + block_size - 1) / block_size; + Mat BT_packed; + BT_packed.create(N * K, (size_t)1u, opt.blob_allocator); + if (BT_packed.empty()) + return -100; + + Mat BT_packed_descales; + BT_packed_descales.create(N * block_count, (size_t)4u, opt.blob_allocator); + if (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(opt.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; + 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); + int kk = 0; + for (; kk + 15 < max_kk; kk += 16) + { + 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; + _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) + { + 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)); + 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; + p0++; + p1++; + p2++; + p3++; + } + + *pd++ = 1.f / *ps0++; + *pd++ = 1.f / *ps1++; + *pd++ = 1.f / *ps2++; + *pd++ = 1.f / *ps3++; + } + } +#endif // __ARM_NEON + #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; + 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); + int kk = 0; +#if __ARM_NEON + for (; kk + 15 < max_kk; kk += 16) + { + int8x16_t _p0 = vld1q_s8(p0); + 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) + { + 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)); +#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; + } +#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 +#endif // __ARM_NEON + for (; kk + 1 < max_kk; kk += 2) + { +#if !__ARM_NEON && __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + // per-k order for the simd32 smlad consumer + pp[0] = p0[0]; + pp[1] = p1[0]; + pp[2] = p0[1]; + pp[3] = p1[1]; +#else + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; +#endif + pp += 4; + p0 += 2; + p1 += 2; + } + if (kk < max_kk) + { + pp[0] = p0[0]; + pp[1] = p1[0]; + pp += 2; + p0++; + p1++; + } + + *pd++ = 1.f / *ps0++; + *pd++ = 1.f / *ps1++; + } + } + #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; + 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); + 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; + } +#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 +#endif // __ARM_NEON + 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++; + + *pd++ = 1.f / *ps0++; + } + } + } + + 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 max_kk, int K, int block_size) +{ +#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, max_kk, 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, max_kk, 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; + const int block_count = (K + block_size - 1) / block_size; + const int block_start = k / block_size; + + 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; + int jj = 0; + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum0; + float32x4_t _fsum1; + float32x4_t _fsum2; + float32x4_t _fsum3; + float32x4_t _fsum4; + float32x4_t _fsum5; + float32x4_t _fsum6; + float32x4_t _fsum7; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + _fsum2 = vdupq_n_f32(0.f); + _fsum3 = vdupq_n_f32(0.f); + _fsum4 = vdupq_n_f32(0.f); + _fsum5 = vdupq_n_f32(0.f); + _fsum6 = vdupq_n_f32(0.f); + _fsum7 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + _fsum2 = vld1q_f32(outptr + 8); + _fsum3 = vld1q_f32(outptr + 12); + _fsum4 = vld1q_f32(outptr + 16); + _fsum5 = vld1q_f32(outptr + 20); + _fsum6 = vld1q_f32(outptr + 24); + _fsum7 = vld1q_f32(outptr + 28); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, 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); + 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_kk0; kk += 8) + { + 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); + _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 + for (; kk + 3 < max_kk0; kk += 4) + { + 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); + _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; + } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM + { + int nn = (max_kk0 - 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_kk0; 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_kk0; kk += 2) + { + 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))))); + _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_kk0) + { + 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))); + 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))); + 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))); + 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; + } + + 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)); + _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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum0; + float32x4_t _fsum1; + float32x4_t _fsum2; + float32x4_t _fsum3; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + _fsum2 = vdupq_n_f32(0.f); + _fsum3 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + _fsum2 = vld1q_f32(outptr + 8); + _fsum3 = vld1q_f32(outptr + 12); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _s0 = vdupq_n_s32(0); + int32x4_t _s1 = vdupq_n_s32(0); + int32x4_t _s2 = vdupq_n_s32(0); + int32x4_t _s3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk0; kk += 8) + { + 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 _b = vld1q_s8(pB); + _s0 = vmmlaq_s32(_s0, _a01, _b); + _s1 = vmmlaq_s32(_s1, _a23, _b); + _s2 = vmmlaq_s32(_s2, _a45, _b); + _s3 = vmmlaq_s32(_s3, _a67, _b); + pA += 64; + pB += 16; + } + int32x4x2_t _ss0 = vuzpq_s32(_s0, _s1); + int32x4x2_t _ss1 = vuzpq_s32(_s2, _s3); + _sum0 = vaddq_s32(_sum0, _ss0.val[0]); + _sum1 = vaddq_s32(_sum1, _ss0.val[1]); + _sum2 = vaddq_s32(_sum2, _ss1.val[0]); + _sum3 = vaddq_s32(_sum3, _ss1.val[1]); +#endif // __ARM_FEATURE_MATMUL_INT8 + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _a1 = vld1q_s8(pA + 16); + int8x8_t _b = vld1_s8(pB); + _sum0 = vdotq_lane_s32(_sum0, _a0, _b, 0); + _sum1 = vdotq_lane_s32(_sum1, _a0, _b, 1); + _sum2 = vdotq_lane_s32(_sum2, _a1, _b, 0); + _sum3 = vdotq_lane_s32(_sum3, _a1, _b, 1); + pA += 32; + pB += 8; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _a1 = vld1q_s8(pA + 16); + int16x4_t _b = vreinterpret_s16_s8(vld1_s8(pB)); + int8x8_t _b0 = vreinterpret_s8_s16(vdup_lane_s16(_b, 0)); + int8x8_t _b1 = vreinterpret_s8_s16(vdup_lane_s16(_b, 1)); + int8x8_t _b2 = vreinterpret_s8_s16(vdup_lane_s16(_b, 2)); + int8x8_t _b3 = vreinterpret_s8_s16(vdup_lane_s16(_b, 3)); + _sum0 = vpadalq_s16(_sum0, vmull_s8(vget_low_s8(_a0), _b0)); + _sum1 = vpadalq_s16(_sum1, vmull_s8(vget_low_s8(_a0), _b1)); + _sum2 = vpadalq_s16(_sum2, vmull_s8(vget_high_s8(_a0), _b0)); + _sum3 = vpadalq_s16(_sum3, vmull_s8(vget_high_s8(_a0), _b1)); + _sum0 = vpadalq_s16(_sum0, vmull_s8(vget_low_s8(_a1), _b2)); + _sum1 = vpadalq_s16(_sum1, vmull_s8(vget_low_s8(_a1), _b3)); + _sum2 = vpadalq_s16(_sum2, vmull_s8(vget_high_s8(_a1), _b2)); + _sum3 = vpadalq_s16(_sum3, vmull_s8(vget_high_s8(_a1), _b3)); + pA += 32; + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + int8x16_t _a = vld1q_s8(pA); + int16x4_t _b = vreinterpret_s16_s32(vld1_dup_s32((const int*)pB)); + int16x4x2_t _b01 = vuzp_s16(_b, _b); + int8x8_t _b0 = vreinterpret_s8_s16(_b01.val[0]); + int8x8_t _b1 = vreinterpret_s8_s16(_b01.val[1]); + int16x8_t _s0 = vmull_s8(vget_low_s8(_a), _b0); + int16x8_t _s1 = vmull_s8(vget_low_s8(_a), _b1); + int16x8_t _s2 = vmull_s8(vget_high_s8(_a), _b0); + int16x8_t _s3 = vmull_s8(vget_high_s8(_a), _b1); + _sum0 = vpadalq_s16(_sum0, _s0); + _sum1 = vpadalq_s16(_sum1, _s1); + _sum2 = vpadalq_s16(_sum2, _s2); + _sum3 = vpadalq_s16(_sum3, _s3); + pA += 16; + pB += 4; + } + if (kk < max_kk0) + { + int8x8_t _a = vld1_s8(pA); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + int8x8x2_t _b01 = vuzp_s8(_b, _b); + int16x8_t _s0 = vmull_s8(_a, _b01.val[0]); + int16x8_t _s1 = vmull_s8(_a, _b01.val[1]); + _sum0 = vaddw_s16(_sum0, vget_low_s16(_s0)); + _sum1 = vaddw_s16(_sum1, vget_low_s16(_s1)); + _sum2 = vaddw_s16(_sum2, vget_high_s16(_s0)); + _sum3 = vaddw_s16(_sum3, vget_high_s16(_s1)); + pA += 8; + pB += 2; + } + + int32x4x2_t _s01 = vzipq_s32(_sum0, _sum1); + int32x4x2_t _s23 = vzipq_s32(_sum2, _sum3); + float32x2_t _bd = vld1_f32(pB_descales); + float32x4_t _bdbd = vcombine_f32(_bd, _bd); + float32x4_t _ad0 = vld1q_f32(pA_descales); + float32x4_t _ad1 = vld1q_f32(pA_descales + 4); + float32x4x2_t _ad01 = vzipq_f32(_ad0, _ad0); + float32x4x2_t _ad23 = vzipq_f32(_ad1, _ad1); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_s01.val[0]), vmulq_f32(_bdbd, _ad01.val[0])); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_s01.val[1]), vmulq_f32(_bdbd, _ad01.val[1])); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_s23.val[0]), vmulq_f32(_bdbd, _ad23.val[0])); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_s23.val[1]), vmulq_f32(_bdbd, _ad23.val[1])); + pA_descales += 8; + pB_descales += 2; + } + + vst1q_f32(outptr, _fsum0); + vst1q_f32(outptr + 4, _fsum1); + vst1q_f32(outptr + 8, _fsum2); + vst1q_f32(outptr + 12, _fsum3); + outptr += 16; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum0; + float32x4_t _fsum1; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; +#if __ARM_FEATURE_DOTPROD + for (; kk + 7 < max_kk0; kk += 8) + { + 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); + int8x8_t _b = vld1_s8(pB); + int8x16_t _bb = vcombine_s8(_b, _b); + int32x4_t _s0 = vdotq_s32(vdupq_n_s32(0), _a01, _bb); + int32x4_t _s1 = vdotq_s32(vdupq_n_s32(0), _a23, _bb); + int32x4_t _s2 = vdotq_s32(vdupq_n_s32(0), _a45, _bb); + int32x4_t _s3 = vdotq_s32(vdupq_n_s32(0), _a67, _bb); + _sum0 = vaddq_s32(_sum0, vpaddq_s32(_s0, _s1)); + _sum1 = vaddq_s32(_sum1, vpaddq_s32(_s2, _s3)); + pA += 64; + pB += 8; + } + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _a1 = vld1q_s8(pA + 16); + int8x16_t _b = vreinterpretq_s8_s32(vld1q_dup_s32((const int*)pB)); + _sum0 = vdotq_s32(_sum0, _a0, _b); + _sum1 = vdotq_s32(_sum1, _a1, _b); + pA += 32; + pB += 4; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _a1 = vld1q_s8(pA + 16); + int16x4_t _b = vreinterpret_s16_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _b0 = vreinterpret_s8_s16(vdup_lane_s16(_b, 0)); + int8x8_t _b1 = vreinterpret_s8_s16(vdup_lane_s16(_b, 1)); + _sum0 = vpadalq_s16(_sum0, vmull_s8(vget_low_s8(_a0), _b0)); + _sum1 = vpadalq_s16(_sum1, vmull_s8(vget_high_s8(_a0), _b0)); + _sum0 = vpadalq_s16(_sum0, vmull_s8(vget_low_s8(_a1), _b1)); + _sum1 = vpadalq_s16(_sum1, vmull_s8(vget_high_s8(_a1), _b1)); + pA += 32; + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + int8x16_t _a = vld1q_s8(pA); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + _sum0 = vpadalq_s16(_sum0, vmull_s8(vget_low_s8(_a), _b)); + _sum1 = vpadalq_s16(_sum1, vmull_s8(vget_high_s8(_a), _b)); + pA += 16; + pB += 2; + } + if (kk < max_kk0) + { + int8x8_t _a = vld1_s8(pA); + int16x8_t _s = vmull_s8(_a, vld1_dup_s8(pB)); + _sum0 = vaddw_s16(_sum0, vget_low_s16(_s)); + _sum1 = vaddw_s16(_sum1, vget_high_s16(_s)); + pA += 8; + pB++; + } + + 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_f32(_bd, _ad0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_f32(_bd, _ad1)); + pA_descales += 8; + pB_descales++; + } + + vst1q_f32(outptr, _fsum0); + vst1q_f32(outptr + 4, _fsum1); + outptr += 8; + pB_panel += K; + pB_descales_panel += block_count; + } + pAT += (size_t)8 * A_hstep; + pAT_descales += (size_t)8 * A_descales_hstep; + } +#endif // __aarch64__ + for (; ii + 3 < max_ii; ii += 4) + { + const signed char* pA0_block = pAT; + const float* pA_descales0_block = pAT_descales; + int jj = 0; + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + float32x4_t _fsum0; + float32x4_t _fsum1; + float32x4_t _fsum2; + float32x4_t _fsum3; + float32x4_t _fsum4; + float32x4_t _fsum5; + float32x4_t _fsum6; + float32x4_t _fsum7; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + _fsum2 = vdupq_n_f32(0.f); + _fsum3 = vdupq_n_f32(0.f); + _fsum4 = vdupq_n_f32(0.f); + _fsum5 = vdupq_n_f32(0.f); + _fsum6 = vdupq_n_f32(0.f); + _fsum7 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + _fsum2 = vld1q_f32(outptr + 8); + _fsum3 = vld1q_f32(outptr + 12); + _fsum4 = vld1q_f32(outptr + 16); + _fsum5 = vld1q_f32(outptr + 20); + _fsum6 = vld1q_f32(outptr + 24); + _fsum7 = vld1q_f32(outptr + 28); + } + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, 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); + 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_kk0; kk += 8) + { + 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); + _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 + for (; kk + 3 < max_kk0; kk += 4) + { + 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); + _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; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; 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); + 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_kk0; kk += 2) + { + 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))))); + _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_kk0) + { + 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))); + _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; + } + + 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)); + _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; + } + + 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; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + float32x4_t _fsum0; + float32x4_t _fsum1; + float32x4_t _fsum2; + float32x4_t _fsum3; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + _fsum2 = vdupq_n_f32(0.f); + _fsum3 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + _fsum2 = vld1q_f32(outptr + 8); + _fsum3 = vld1q_f32(outptr + 12); + } + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, 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); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk0; kk += 8) + { + 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); + _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 + for (; kk + 3 < max_kk0; kk += 4) + { + 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); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a, 3); + pA += 16; + pB += 16; + } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM && !__aarch64__ + { + int nn = (max_kk0 - 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_kk0; 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); + 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_kk0; kk += 2) + { + 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))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))))); + pA += 8; + pB += 8; + } + if (kk < max_kk0) + { + 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))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + pA += 4; + pB += 4; + } + + 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))); + _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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + float32x4_t _fsum0; + float32x4_t _fsum1; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _b = vld1q_s8(pB); + _sum0 = vmmlaq_s32(_sum0, _a01, _b); + _sum1 = vmmlaq_s32(_sum1, _a23, _b); + pA += 32; + pB += 16; + } +#endif // __ARM_FEATURE_MATMUL_INT8 + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _a = vld1q_s8(pA); + int8x8_t _b = vld1_s8(pB); + int32x4_t _s0 = vdotq_lane_s32(vdupq_n_s32(0), _a, _b, 0); + int32x4_t _s1 = vdotq_lane_s32(vdupq_n_s32(0), _a, _b, 1); + int32x4x2_t _s01 = vzipq_s32(_s0, _s1); + _sum0 = vaddq_s32(_sum0, _s01.val[0]); + _sum1 = vaddq_s32(_sum1, _s01.val[1]); + pA += 16; + pB += 8; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int8x8_t _a0 = vld1_s8(pA); + int8x8_t _a1 = vld1_s8(pA + 8); + int16x4_t _b = vreinterpret_s16_s8(vld1_s8(pB)); + int8x8_t _b0 = vreinterpret_s8_s16(vdup_lane_s16(_b, 0)); + int8x8_t _b1 = vreinterpret_s8_s16(vdup_lane_s16(_b, 1)); + int8x8_t _b2 = vreinterpret_s8_s16(vdup_lane_s16(_b, 2)); + int8x8_t _b3 = vreinterpret_s8_s16(vdup_lane_s16(_b, 3)); + int32x4_t _s0 = vpaddlq_s16(vmull_s8(_a0, _b0)); + int32x4_t _s1 = vpaddlq_s16(vmull_s8(_a0, _b1)); + _s0 = vpadalq_s16(_s0, vmull_s8(_a1, _b2)); + _s1 = vpadalq_s16(_s1, vmull_s8(_a1, _b3)); + int32x4x2_t _s01 = vzipq_s32(_s0, _s1); + _sum0 = vaddq_s32(_sum0, _s01.val[0]); + _sum1 = vaddq_s32(_sum1, _s01.val[1]); + pA += 16; + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + int8x8_t _a = vld1_s8(pA); + int16x4_t _b = vreinterpret_s16_s32(vld1_dup_s32((const int*)pB)); + int16x4x2_t _b01 = vuzp_s16(_b, _b); + int32x4_t _s0 = vpaddlq_s16(vmull_s8(_a, vreinterpret_s8_s16(_b01.val[0]))); + int32x4_t _s1 = vpaddlq_s16(vmull_s8(_a, vreinterpret_s8_s16(_b01.val[1]))); + int32x4x2_t _s01 = vzipq_s32(_s0, _s1); + _sum0 = vaddq_s32(_sum0, _s01.val[0]); + _sum1 = vaddq_s32(_sum1, _s01.val[1]); + pA += 8; + pB += 4; + } + if (kk < max_kk0) + { + int8x8_t _a = vld1_s8(pA); + int32x4_t _s0 = vmovl_s16(vget_low_s16(vmull_s8(_a, vdup_n_s8(pB[0])))); + int32x4_t _s1 = vmovl_s16(vget_low_s16(vmull_s8(_a, vdup_n_s8(pB[1])))); + int32x4x2_t _s01 = vzipq_s32(_s0, _s1); + _sum0 = vaddq_s32(_sum0, _s01.val[0]); + _sum1 = vaddq_s32(_sum1, _s01.val[1]); + pA += 4; + pB += 2; + } + + float32x2_t _bd = vld1_f32(pB_descales); + float32x4_t _bdbd = vcombine_f32(_bd, _bd); + float32x4_t _ad = vld1q_f32(pA_descales); + float32x4x2_t _ad01 = vzipq_f32(_ad, _ad); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_f32(_bdbd, _ad01.val[0])); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_f32(_bdbd, _ad01.val[1])); + pA_descales += 4; + pB_descales += 2; + } + + vst1q_f32(outptr, _fsum0); + vst1q_f32(outptr + 4, _fsum1); + outptr += 8; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + float32x4_t _fsum; + + if (k == 0) + _fsum = vdupq_n_f32(0.f); + else + _fsum = vld1q_f32(outptr); + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; +#if __ARM_FEATURE_DOTPROD + for (; kk + 7 < max_kk0; kk += 8) + { + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); + int8x8_t _b = vld1_s8(pB); + int8x16_t _bb = vcombine_s8(_b, _b); + int32x4_t _s0 = vdotq_s32(vdupq_n_s32(0), _a01, _bb); + int32x4_t _s1 = vdotq_s32(vdupq_n_s32(0), _a23, _bb); + _sum = vaddq_s32(_sum, vcombine_s32(vpadd_s32(vget_low_s32(_s0), vget_high_s32(_s0)), vpadd_s32(vget_low_s32(_s1), vget_high_s32(_s1)))); + pA += 32; + pB += 8; + } + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _a = vld1q_s8(pA); + int8x16_t _b = vreinterpretq_s8_s32(vld1q_dup_s32((const int*)pB)); + _sum = vdotq_s32(_sum, _a, _b); + pA += 16; + pB += 4; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int8x8_t _a0 = vld1_s8(pA); + int8x8_t _a1 = vld1_s8(pA + 8); + int16x4_t _b = vreinterpret_s16_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _b0 = vreinterpret_s8_s16(vdup_lane_s16(_b, 0)); + int8x8_t _b1 = vreinterpret_s8_s16(vdup_lane_s16(_b, 1)); + _sum = vpadalq_s16(_sum, vmull_s8(_a0, _b0)); + _sum = vpadalq_s16(_sum, vmull_s8(_a1, _b1)); + pA += 16; + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + int8x8_t _a = vld1_s8(pA); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + _sum = vpadalq_s16(_sum, vmull_s8(_a, _b)); + pA += 8; + pB += 2; + } + if (kk < max_kk0) + { + int8x8_t _a = vld1_s8(pA); + int16x8_t _s = vmull_s8(_a, vld1_dup_s8(pB)); + _sum = vaddw_s16(_sum, vget_low_s16(_s)); + pA += 4; + pB++; + } + + float32x4_t _bd = vdupq_n_f32(pB_descales[0]); + float32x4_t _ad = vld1q_f32(pA_descales); + _fsum = vmlaq_f32(_fsum, vcvtq_f32_s32(_sum), vmulq_f32(_bd, _ad)); + pA_descales += 4; + pB_descales++; + } + + vst1q_f32(outptr, _fsum); + outptr += 4; + pB_panel += K; + pB_descales_panel += block_count; + } + pAT += (size_t)4 * A_hstep; + pAT_descales += (size_t)4 * A_descales_hstep; + } +#endif // __ARM_NEON + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pA_block = pAT; + const float* pA_descales_block = pAT_descales; + int jj = 0; + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; +#if __ARM_NEON +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + float32x4_t _fsum0; + float32x4_t _fsum1; + float32x4_t _fsum2; + float32x4_t _fsum3; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + _fsum2 = vdupq_n_f32(0.f); + _fsum3 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + _fsum2 = vld1q_f32(outptr + 8); + _fsum3 = vld1q_f32(outptr + 12); + } + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, 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); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk0; kk += 8) + { + 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); + _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 + for (; kk + 3 < max_kk0; kk += 4) + { + 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); + _sum3 = vdotq_laneq_s32(_sum3, _b1, _a, 1); + pA += 8; + pB0 += 16; + pB1 += 16; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + 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_kk0; kk += 2) + { + 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))))); + _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_kk0) + { + 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))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_high_s16(_p1))); + pA += 2; + pB0 += 4; + pB1 += 4; + } + + 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)); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_lane_f32(_bd1, _ad, 1)); + + pA_descales += 2; + pB_descales0 += 4; + pB_descales1 += 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; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + float32x4_t _fsum0; + float32x4_t _fsum1; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + } + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; kk += 8) + { + 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; + 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 + for (; kk + 3 < max_kk0; kk += 4) + { + 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_kk0; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + 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_kk0; kk += 2) + { + 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; + pB += 8; + } + if (kk < max_kk0) + { + 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; + } + + 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)); + + pA_descales += 2; + pB_descales += 4; + } + + vst1q_f32(outptr, _fsum0); + outptr += 4; + vst1q_f32(outptr, _fsum1); + outptr += 4; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // __ARM_NEON + for (; jj + 1 < max_jj; jj += 2) + { +#if __ARM_NEON + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum; + + if (k == 0) + _fsum = vdupq_n_f32(0.f); + else + _fsum = vld1q_f32(outptr); + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk0; kk += 8) + { + int8x16_t _a = vld1q_s8(pA); + int8x16_t _b = vld1q_s8(pB); + _sum = vmmlaq_s32(_sum, _a, _b); + pA += 16; + pB += 16; + } +#endif // __ARM_FEATURE_MATMUL_INT8 + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + int8x8_t _b = vld1_s8(pB); + int32x4_t _s0 = vdotq_lane_s32(vdupq_n_s32(0), _a, _b, 0); + int32x4_t _s1 = vdotq_lane_s32(vdupq_n_s32(0), _a, _b, 1); + _sum = vaddq_s32(_sum, vzipq_s32(_s0, _s1).val[0]); + pA += 8; + pB += 8; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int8x8_t _a0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pA)); + int8x8_t _a1 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pA + 4))); + int16x4_t _b = vreinterpret_s16_s8(vld1_s8(pB)); + int8x8_t _b0 = vreinterpret_s8_s16(vdup_lane_s16(_b, 0)); + int8x8_t _b1 = vreinterpret_s8_s16(vdup_lane_s16(_b, 1)); + int8x8_t _b2 = vreinterpret_s8_s16(vdup_lane_s16(_b, 2)); + int8x8_t _b3 = vreinterpret_s8_s16(vdup_lane_s16(_b, 3)); + int32x4_t _s0 = vpaddlq_s16(vmull_s8(_a0, _b0)); + int32x4_t _s1 = vpaddlq_s16(vmull_s8(_a0, _b1)); + _s0 = vpadalq_s16(_s0, vmull_s8(_a1, _b2)); + _s1 = vpadalq_s16(_s1, vmull_s8(_a1, _b3)); + _sum = vaddq_s32(_sum, vzipq_s32(_s0, _s1).val[0]); + pA += 8; + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + int8x8_t _a = vreinterpret_s8_s32(vld1_dup_s32((const int*)pA)); + int16x4_t _b = vreinterpret_s16_s32(vld1_dup_s32((const int*)pB)); + int16x4x2_t _b01 = vuzp_s16(_b, _b); + int32x4_t _s0 = vpaddlq_s16(vmull_s8(_a, vreinterpret_s8_s16(_b01.val[0]))); + int32x4_t _s1 = vpaddlq_s16(vmull_s8(_a, vreinterpret_s8_s16(_b01.val[1]))); + _sum = vaddq_s32(_sum, vzipq_s32(_s0, _s1).val[0]); + pA += 4; + pB += 4; + } + if (kk < max_kk0) + { + int8x8_t _a = vreinterpret_s8_s16(vld1_dup_s16((const short*)pA)); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + int8x8_t _aa = vzip_s8(_a, _a).val[0]; + _sum = vaddq_s32(_sum, vmovl_s16(vget_low_s16(vmull_s8(_aa, _b)))); + pA += 2; + pB += 2; + } + + float32x2_t _ad = vld1_f32(pA_descales); + float32x2_t _bd = vld1_f32(pB_descales); + float32x4_t _adad = vcombine_f32(_ad, _ad); + float32x4_t _bdbd = vcombine_f32(_bd, _bd); + _fsum = vmlaq_f32(_fsum, vcvtq_f32_s32(_sum), vmulq_f32(vzipq_f32(_adad, _adad).val[0], _bdbd)); + pA_descales += 2; + pB_descales += 2; + } + + vst1q_f32(outptr, _fsum); + outptr += 4; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; +#elif __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float fsum00; + float fsum01; + float fsum10; + float fsum11; + + if (k == 0) + { + fsum00 = 0.f; + fsum01 = 0.f; + fsum10 = 0.f; + fsum11 = 0.f; + } + else + { + fsum00 = outptr[0]; + fsum01 = outptr[1]; + fsum10 = outptr[2]; + fsum11 = outptr[3]; + } + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum00 = 0; + int sum01 = 0; + int sum10 = 0; + int sum11 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 1 < max_kk0; 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_kk0) + { + 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; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; +#else + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float fsum00; + float fsum01; + float fsum10; + float fsum11; + + if (k == 0) + { + fsum00 = 0.f; + fsum01 = 0.f; + fsum10 = 0.f; + fsum11 = 0.f; + } + else + { + fsum00 = outptr[0]; + fsum01 = outptr[1]; + fsum10 = outptr[2]; + fsum11 = outptr[3]; + } + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum00 = 0; + int sum01 = 0; + int sum10 = 0; + int sum11 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 1 < max_kk0; kk += 2) + { + const int b00 = pB[0]; + const int b01 = pB[1]; + const int b10 = pB[2]; + const int b11 = pB[3]; + sum00 += pA[0] * b00 + pA[2] * b01; + sum01 += pA[0] * b10 + pA[2] * b11; + sum10 += pA[1] * b00 + pA[3] * b01; + sum11 += pA[1] * b10 + pA[3] * b11; + pA += 4; + pB += 4; + } + if (kk < max_kk0) + { + const int b0 = pB[0]; + const int b1 = pB[1]; + sum00 += pA[0] * b0; + sum01 += pA[0] * b1; + sum10 += pA[1] * b0; + sum11 += pA[1] * b1; + pA += 2; + pB += 2; + } + + const float bd0 = pB_descales[0]; + const float bd1 = pB_descales[1]; + const float ad0 = pA_descales[0]; + fsum00 += sum00 * ad0 * bd0; + fsum01 += sum01 * ad0 * bd1; + const float ad1 = pA_descales[1]; + fsum10 += sum10 * ad1 * bd0; + fsum11 += sum11 * ad1 * bd1; + + pA_descales += 2; + pB_descales += 2; + } + + outptr[0] = fsum00; + outptr++; + outptr[0] = fsum01; + outptr++; + outptr[0] = fsum10; + outptr++; + outptr[0] = fsum11; + outptr++; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; +#endif // __ARM_NEON + } + for (; jj < max_jj; jj++) + { +#if __ARM_NEON + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum; + + if (k == 0) + _fsum = vdupq_n_f32(0.f); + else + _fsum = vcombine_f32(vld1_f32(outptr), vdup_n_f32(0.f)); + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; +#if __ARM_FEATURE_DOTPROD + for (; kk + 7 < max_kk0; kk += 8) + { + int8x16_t _a = vld1q_s8(pA); + int8x8_t _b = vld1_s8(pB); + int8x16_t _bb = vcombine_s8(_b, _b); + int32x4_t _s = vdotq_s32(vdupq_n_s32(0), _a, _bb); + _sum = vaddq_s32(_sum, vpaddq_s32(_s, _s)); + pA += 16; + pB += 8; + } + for (; kk + 3 < max_kk0; kk += 4) + { + int8x8_t _a = vld1_s8(pA); + int8x16_t _aa = vcombine_s8(_a, vdup_n_s8(0)); + int8x16_t _b = vreinterpretq_s8_s32(vld1q_dup_s32((const int*)pB)); + _sum = vdotq_s32(_sum, _aa, _b); + pA += 8; + pB += 4; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int8x8_t _a0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pA)); + int8x8_t _a1 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pA + 4))); + int16x4_t _b = vreinterpret_s16_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _b0 = vreinterpret_s8_s16(vdup_lane_s16(_b, 0)); + int8x8_t _b1 = vreinterpret_s8_s16(vdup_lane_s16(_b, 1)); + _sum = vaddq_s32(_sum, vpaddlq_s16(vmull_s8(_a0, _b0))); + _sum = vaddq_s32(_sum, vpaddlq_s16(vmull_s8(_a1, _b1))); + pA += 8; + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + int8x8_t _a = vreinterpret_s8_s32(vld1_dup_s32((const int*)pA)); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + _sum = vaddq_s32(_sum, vpaddlq_s16(vmull_s8(_a, _b))); + pA += 4; + pB += 2; + } + if (kk < max_kk0) + { + int8x8_t _a = vreinterpret_s8_s16(vld1_dup_s16((const short*)pA)); + int16x8_t _s = vmull_s8(_a, vld1_dup_s8(pB)); + _sum = vaddw_s16(_sum, vget_low_s16(_s)); + pA += 2; + pB++; + } + + float32x2_t _ad = vld1_f32(pA_descales); + float32x4_t _scale = vmulq_n_f32(vcombine_f32(_ad, _ad), pB_descales[0]); + _fsum = vmlaq_f32(_fsum, vcvtq_f32_s32(_sum), _scale); + pA_descales += 2; + pB_descales++; + } + + vst1_f32(outptr, vget_low_f32(_fsum)); + outptr += 2; + pB_panel += K; + pB_descales_panel += block_count; +#elif __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float fsum00; + float fsum10; + + if (k == 0) + { + fsum00 = 0.f; + fsum10 = 0.f; + } + else + { + fsum00 = outptr[0]; + fsum10 = outptr[1]; + } + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum00 = 0; + int sum10 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + for (int kk = 0; kk < max_kk0; 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; + pB_panel += K; + pB_descales_panel += block_count; +#else + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float fsum00; + float fsum10; + + if (k == 0) + { + fsum00 = 0.f; + fsum10 = 0.f; + } + else + { + fsum00 = outptr[0]; + fsum10 = outptr[1]; + } + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum00 = 0; + int sum10 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 1 < max_kk0; kk += 2) + { + const int b00 = pB[0]; + const int b01 = pB[1]; + sum00 += pA[0] * b00 + pA[2] * b01; + sum10 += pA[1] * b00 + pA[3] * b01; + pA += 4; + pB += 2; + } + if (kk < max_kk0) + { + const int b0 = pB[0]; + sum00 += pA[0] * b0; + sum10 += pA[1] * b0; + pA += 2; + pB += 1; + } + + const float bd0 = pB_descales[0]; + const float ad0 = pA_descales[0]; + fsum00 += sum00 * ad0 * bd0; + const float ad1 = pA_descales[1]; + fsum10 += sum10 * ad1 * bd0; + + pA_descales += 2; + pB_descales += 1; + } + + outptr[0] = fsum00; + outptr++; + outptr[0] = fsum10; + outptr++; + pB_panel += K; + pB_descales_panel += block_count; +#endif // __ARM_NEON + } + pAT += (size_t)2 * A_hstep; + pAT_descales += (size_t)2 * A_descales_hstep; + } + for (; ii < max_ii; ii++) + { + const signed char* pA0_block = pAT; + const float* pA_descales0_block = pAT_descales; + int jj = 0; + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; +#if __ARM_NEON +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + float32x4_t _fsum0; + float32x4_t _fsum1; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + _fsum1 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + _fsum1 = vld1q_f32(outptr + 4); + } + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + const signed char* pA = pA0; + 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); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk0; kk += 8) + { + 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(pA), 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); + pA += 8; + 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 + for (; kk + 3 < max_kk0; kk += 4) + { + 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*)pA, vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + _sum1 = vdotq_s32(_sum1, _b1, _a0); + pA += 4; + pB0 += 16; + pB1 += 16; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, 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); + pA += 4; + pB0 += 16; + pB1 += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + 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*)pA, 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))); + pA += 2; + pB0 += 8; + pB1 += 8; + } + if (kk < max_kk0) + { + 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(pA, 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))); + pA++; + pB0 += 4; + pB1 += 4; + } + + 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)); + + pA0 = pA; + pA_descales0++; + pB_descales0 += 4; + pB_descales1 += 4; + } + + vst1q_f32(outptr, _fsum0); + outptr += 4; + vst1q_f32(outptr, _fsum1); + outptr += 4; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + float32x4_t _fsum0; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vld1q_f32(outptr); + } + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + const signed char* pA = pA0; + 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_kk0; kk += 8) + { + int8x16_t _b0 = vld1q_s8(pB + 0); + int8x16_t _b1 = vld1q_s8(pB + 16); + int8x16_t _a0 = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + _msum1 = vmmlaq_s32(_msum1, _a0, _b1); + pA += 8; + pB += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); +#endif // __ARM_FEATURE_MATMUL_INT8 + for (; kk + 3 < max_kk0; kk += 4) + { + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + pA += 4; + pB += 16; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, 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); + pA += 4; + pB += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + int8x8_t _b0 = vld1_s8(pB); + int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); + pA += 2; + pB += 8; + } + if (kk < max_kk0) + { + int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + int8x8_t _a0 = vld1_lane_s8(pA, 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))); + pA++; + pB += 4; + } + + 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 = pA; + pA_descales0++; + pB_descales += 4; + } + + vst1q_f32(outptr, _fsum0); + outptr += 4; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // __ARM_NEON + for (; jj + 1 < max_jj; jj += 2) + { +#if __ARM_NEON + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float32x4_t _fsum0; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vcombine_f32(vld1_f32(outptr), vdup_n_f32(0.f)); + } + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + const signed char* pA = pA0; + int kk = 0; +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + for (; kk + 7 < max_kk0; kk += 8) + { + int8x16_t _b0 = vld1q_s8(pB + 0); + int8x16_t _a0 = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + pA += 8; + pB += 16; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vdup_n_s32(0)); +#endif // __ARM_FEATURE_MATMUL_INT8 + for (; kk + 3 < max_kk0; kk += 4) + { + 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*)pA, vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + pA += 4; + pB += 8; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, 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); + pA += 4; + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + 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*)pA, vdup_n_s16(0), 0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); + pA += 2; + pB += 4; + } + if (kk < max_kk0) + { + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + int8x8_t _a0 = vld1_lane_s8(pA, 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))); + pA++; + pB += 2; + } + + 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 = pA; + pA_descales0++; + pB_descales += 2; + } + + vst1_f32(outptr, vget_low_f32(_fsum0)); + outptr += 2; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; +#else // __ARM_NEON + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float fsum00; + float fsum01; + + if (k == 0) + { + fsum00 = 0.f; + fsum01 = 0.f; + } + else + { + fsum00 = outptr[0]; + fsum01 = outptr[1]; + } + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum00 = 0; + int sum01 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 1 < max_kk0; 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[0] * b00 + pA0[1] * b01; + sum01 += pA0[0] * b10 + pA0[1] * b11; + pA0 += 2; + pB += 4; + } + if (kk < max_kk0) + { + const int b0 = pB[0]; + const int b1 = pB[1]; + sum00 += pA0[0] * b0; + sum01 += pA0[0] * b1; + pA0++; + 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; + + pA_descales0++; + pB_descales += 2; + } + + outptr[0] = fsum00; + outptr++; + outptr[0] = fsum01; + outptr++; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; +#endif // __ARM_NEON + } + for (; jj < max_jj; jj++) + { +#if __ARM_NEON + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float32x4_t _fsum0; + + if (k == 0) + { + _fsum0 = vdupq_n_f32(0.f); + } + else + { + _fsum0 = vsetq_lane_f32(outptr[0], vdupq_n_f32(0.f), 0); + } + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + const signed char* pA = pA0; + int kk = 0; +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + 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*)pA, vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + pA += 4; + pB += 4; + } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk0; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, 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); + pA += 4; + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk0; kk += 2) + { + 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*)pA, vdup_n_s16(0), 0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); + pA += 2; + pB += 2; + } + if (kk < max_kk0) + { + int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + int8x8_t _a0 = vld1_lane_s8(pA, 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))); + pA++; + pB += 1; + } + + 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 = pA; + pA_descales0++; + pB_descales += 1; + } + + vst1q_lane_f32(outptr, _fsum0, 0); + outptr++; + pB_panel += K; + pB_descales_panel += block_count; +#else // __ARM_NEON + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float fsum00; + + if (k == 0) + { + fsum00 = 0.f; + } + else + { + fsum00 = outptr[0]; + } + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum00 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 1 < max_kk0; kk += 2) + { + const int b00 = pB[0]; + const int b01 = pB[1]; + sum00 += pA0[0] * b00 + pA0[1] * b01; + pA0 += 2; + pB += 2; + } + if (kk < max_kk0) + { + const int b0 = pB[0]; + sum00 += pA0[0] * b0; + pA0++; + pB += 1; + } + + const float bd0 = pB_descales[0]; + const float ad0 = pA_descales0[0]; + fsum00 += sum00 * ad0 * bd0; + + pA_descales0++; + pB_descales += 1; + } + + outptr[0] = fsum00; + outptr++; + pB_panel += K; + pB_descales_panel += block_count; +#endif // __ARM_NEON + } + pAT += A_hstep; + pAT_descales += A_descales_hstep; + } +} + +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(); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + + int tile_size = (int)sqrtf((float)l2_cache_size / (2 * sizeof(signed char) + sizeof(float))); + +#if __aarch64__ + TILE_M = std::max(8, tile_size / 8 * 8); + TILE_N = std::max(8, tile_size / 8 * 8); +#elif __ARM_NEON + TILE_M = std::max(4, tile_size / 4 * 4); + TILE_N = std::max(4, tile_size / 4 * 4); +#else + TILE_M = std::max(2, tile_size / 2 * 2); + TILE_N = std::max(2, tile_size / 2 * 2); +#endif + + TILE_K = std::max(block_size, tile_size / block_size * block_size); + + if (K > 0) + { + int nn_K = (K + TILE_K - 1) / TILE_K; + TILE_K = std::min(TILE_K, ((K + nn_K - 1) / nn_K + block_size - 1) / block_size * block_size); + TILE_K = std::min(TILE_K, K); + + if (nn_K == 1) + { + tile_size = (int)((float)l2_cache_size / 2 / sizeof(signed char) / TILE_K); + +#if __aarch64__ + TILE_M = std::max(8, tile_size / 8 * 8); + TILE_N = std::max(8, tile_size / 8 * 8); +#elif __ARM_NEON + TILE_M = std::max(4, tile_size / 4 * 4); + TILE_N = std::max(4, tile_size / 4 * 4); +#else + TILE_M = std::max(2, tile_size / 2 * 2); + TILE_N = std::max(2, tile_size / 2 * 2); +#endif + } + } + + TILE_M *= std::min(nT, get_physical_cpu_count()); + + if (M > 0) + { + int nn_M = (M + TILE_M - 1) / TILE_M; +#if __aarch64__ + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + 7) / 8 * 8); +#elif __ARM_NEON + 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 + } + + if (N > 0) + { + 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 + } + + if (nT > 1) + { +#if __aarch64__ + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + 7) / 8 * 8); +#elif __ARM_NEON + 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 + } + + // 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 + } + + 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); + } +} + +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, 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)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)(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 + 7 < max_jj; jj += 8) + { + float32x4_t _out00 = vld1q_f32(pp); + float32x4_t _out01 = vld1q_f32(pp + 32); + float32x4_t _out10 = vld1q_f32(pp + 4); + float32x4_t _out11 = vld1q_f32(pp + 36); + float32x4_t _out20 = vld1q_f32(pp + 8); + float32x4_t _out21 = vld1q_f32(pp + 40); + float32x4_t _out30 = vld1q_f32(pp + 12); + float32x4_t _out31 = vld1q_f32(pp + 44); + + float32x4_t _out40 = vld1q_f32(pp + 16); + float32x4_t _out41 = vld1q_f32(pp + 48); + float32x4_t _out50 = vld1q_f32(pp + 20); + float32x4_t _out51 = vld1q_f32(pp + 52); + float32x4_t _out60 = vld1q_f32(pp + 24); + float32x4_t _out61 = vld1q_f32(pp + 56); + float32x4_t _out70 = vld1q_f32(pp + 28); + float32x4_t _out71 = vld1q_f32(pp + 60); + pp += 64; + + if (pC) + { + if (broadcast_type_C <= 2) + { + float32x4_t _c0123; + float32x4_t _c4567; + if (broadcast_type_C == 0) + { + _c0123 = vdupq_n_f32(c0); + _c4567 = _c0123; + } + else + { + _c0123 = vmulq_n_f32(vld1q_f32(pC), beta); + _c4567 = vmulq_n_f32(vld1q_f32(pC + 4), beta); + } + _out00 = vaddq_f32(_out00, vdupq_laneq_f32(_c0123, 0)); + _out01 = vaddq_f32(_out01, vdupq_laneq_f32(_c0123, 0)); + _out10 = vaddq_f32(_out10, vdupq_laneq_f32(_c0123, 1)); + _out11 = vaddq_f32(_out11, vdupq_laneq_f32(_c0123, 1)); + _out20 = vaddq_f32(_out20, vdupq_laneq_f32(_c0123, 2)); + _out21 = vaddq_f32(_out21, vdupq_laneq_f32(_c0123, 2)); + _out30 = vaddq_f32(_out30, vdupq_laneq_f32(_c0123, 3)); + _out31 = vaddq_f32(_out31, vdupq_laneq_f32(_c0123, 3)); + _out40 = vaddq_f32(_out40, vdupq_laneq_f32(_c4567, 0)); + _out41 = vaddq_f32(_out41, vdupq_laneq_f32(_c4567, 0)); + _out50 = vaddq_f32(_out50, vdupq_laneq_f32(_c4567, 1)); + _out51 = vaddq_f32(_out51, vdupq_laneq_f32(_c4567, 1)); + _out60 = vaddq_f32(_out60, vdupq_laneq_f32(_c4567, 2)); + _out61 = vaddq_f32(_out61, vdupq_laneq_f32(_c4567, 2)); + _out70 = vaddq_f32(_out70, vdupq_laneq_f32(_c4567, 3)); + _out71 = vaddq_f32(_out71, vdupq_laneq_f32(_c4567, 3)); + } + if (broadcast_type_C == 3) + { + float32x4_t _c0 = vld1q_f32(pC); + float32x4_t _c1 = vld1q_f32(pC + 4); + float32x4_t _c2 = vld1q_f32(pC + c_hstep); + float32x4_t _c3 = vld1q_f32(pC + c_hstep + 4); + float32x4_t _c4 = vld1q_f32(pC + c_hstep * 2); + float32x4_t _c5 = vld1q_f32(pC + c_hstep * 2 + 4); + float32x4_t _c6 = vld1q_f32(pC + c_hstep * 3); + float32x4_t _c7 = vld1q_f32(pC + c_hstep * 3 + 4); + if (beta == 1.f) + { + _out00 = vaddq_f32(_out00, _c0); + _out01 = vaddq_f32(_out01, _c1); + _out10 = vaddq_f32(_out10, _c2); + _out11 = vaddq_f32(_out11, _c3); + _out20 = vaddq_f32(_out20, _c4); + _out21 = vaddq_f32(_out21, _c5); + _out30 = vaddq_f32(_out30, _c6); + _out31 = vaddq_f32(_out31, _c7); + } + else + { + _out00 = vmlaq_n_f32(_out00, _c0, beta); + _out01 = vmlaq_n_f32(_out01, _c1, beta); + _out10 = vmlaq_n_f32(_out10, _c2, beta); + _out11 = vmlaq_n_f32(_out11, _c3, beta); + _out20 = vmlaq_n_f32(_out20, _c4, beta); + _out21 = vmlaq_n_f32(_out21, _c5, beta); + _out30 = vmlaq_n_f32(_out30, _c6, beta); + _out31 = vmlaq_n_f32(_out31, _c7, beta); + } + _c0 = vld1q_f32(pC + c_hstep * 4); + _c1 = vld1q_f32(pC + c_hstep * 4 + 4); + _c2 = vld1q_f32(pC + c_hstep * 5); + _c3 = vld1q_f32(pC + c_hstep * 5 + 4); + _c4 = vld1q_f32(pC + c_hstep * 6); + _c5 = vld1q_f32(pC + c_hstep * 6 + 4); + _c6 = vld1q_f32(pC + c_hstep * 7); + _c7 = vld1q_f32(pC + c_hstep * 7 + 4); + if (beta == 1.f) + { + _out40 = vaddq_f32(_out40, _c0); + _out41 = vaddq_f32(_out41, _c1); + _out50 = vaddq_f32(_out50, _c2); + _out51 = vaddq_f32(_out51, _c3); + _out60 = vaddq_f32(_out60, _c4); + _out61 = vaddq_f32(_out61, _c5); + _out70 = vaddq_f32(_out70, _c6); + _out71 = vaddq_f32(_out71, _c7); + } + else + { + _out40 = vmlaq_n_f32(_out40, _c0, beta); + _out41 = vmlaq_n_f32(_out41, _c1, beta); + _out50 = vmlaq_n_f32(_out50, _c2, beta); + _out51 = vmlaq_n_f32(_out51, _c3, beta); + _out60 = vmlaq_n_f32(_out60, _c4, beta); + _out61 = vmlaq_n_f32(_out61, _c5, beta); + _out70 = vmlaq_n_f32(_out70, _c6, beta); + _out71 = vmlaq_n_f32(_out71, _c7, beta); + } + pC += 8; + } + if (broadcast_type_C == 4) + { + float32x4_t _c = vld1q_f32(pC); + if (beta != 1.f) + _c = vmulq_n_f32(_c, beta); + _out00 = vaddq_f32(_out00, _c); + _out10 = vaddq_f32(_out10, _c); + _out20 = vaddq_f32(_out20, _c); + _out30 = vaddq_f32(_out30, _c); + _out40 = vaddq_f32(_out40, _c); + _out50 = vaddq_f32(_out50, _c); + _out60 = vaddq_f32(_out60, _c); + _out70 = vaddq_f32(_out70, _c); + _c = vld1q_f32(pC + 4); + if (beta != 1.f) + _c = vmulq_n_f32(_c, beta); + _out01 = vaddq_f32(_out01, _c); + _out11 = vaddq_f32(_out11, _c); + _out21 = vaddq_f32(_out21, _c); + _out31 = vaddq_f32(_out31, _c); + _out41 = vaddq_f32(_out41, _c); + _out51 = vaddq_f32(_out51, _c); + _out61 = vaddq_f32(_out61, _c); + _out71 = vaddq_f32(_out71, _c); + pC += 8; + } + } + if (alpha != 1.f) + { + _out00 = vmulq_n_f32(_out00, alpha); + _out01 = vmulq_n_f32(_out01, alpha); + _out10 = vmulq_n_f32(_out10, alpha); + _out11 = vmulq_n_f32(_out11, alpha); + _out20 = vmulq_n_f32(_out20, alpha); + _out21 = vmulq_n_f32(_out21, alpha); + _out30 = vmulq_n_f32(_out30, alpha); + _out31 = vmulq_n_f32(_out31, alpha); + _out40 = vmulq_n_f32(_out40, alpha); + _out41 = vmulq_n_f32(_out41, alpha); + _out50 = vmulq_n_f32(_out50, alpha); + _out51 = vmulq_n_f32(_out51, alpha); + _out60 = vmulq_n_f32(_out60, alpha); + _out61 = vmulq_n_f32(_out61, alpha); + _out70 = vmulq_n_f32(_out70, alpha); + _out71 = vmulq_n_f32(_out71, alpha); + } + + vst1q_f32(outptr0, _out00); + vst1q_f32(outptr0 + 4, _out01); + vst1q_f32(outptr1, _out10); + vst1q_f32(outptr1 + 4, _out11); + vst1q_f32(outptr2, _out20); + vst1q_f32(outptr2 + 4, _out21); + vst1q_f32(outptr3, _out30); + vst1q_f32(outptr3 + 4, _out31); + + vst1q_f32(outptr4, _out40); + vst1q_f32(outptr4 + 4, _out41); + vst1q_f32(outptr5, _out50); + vst1q_f32(outptr5 + 4, _out51); + vst1q_f32(outptr6, _out60); + vst1q_f32(outptr6 + 4, _out61); + vst1q_f32(outptr7, _out70); + vst1q_f32(outptr7 + 4, _out71); + + outptr0 += 8; + outptr1 += 8; + outptr2 += 8; + outptr3 += 8; + outptr4 += 8; + outptr5 += 8; + outptr6 += 8; + outptr7 += 8; + } + float32x4_t _c0123 = vdupq_n_f32(c0); + _c0123 = vsetq_lane_f32(c1, _c0123, 1); + _c0123 = vsetq_lane_f32(c2, _c0123, 2); + _c0123 = vsetq_lane_f32(c3, _c0123, 3); + float32x4_t _c4567 = vdupq_n_f32(c4); + _c4567 = vsetq_lane_f32(c5, _c4567, 1); + _c4567 = vsetq_lane_f32(c6, _c4567, 2); + _c4567 = vsetq_lane_f32(c7, _c4567, 3); + 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); + pp += 32; + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, vdupq_laneq_f32(_c0123, 0)); + _out1 = vaddq_f32(_out1, vdupq_laneq_f32(_c0123, 1)); + _out2 = vaddq_f32(_out2, vdupq_laneq_f32(_c0123, 2)); + _out3 = vaddq_f32(_out3, vdupq_laneq_f32(_c0123, 3)); + _out4 = vaddq_f32(_out4, vdupq_laneq_f32(_c4567, 0)); + _out5 = vaddq_f32(_out5, vdupq_laneq_f32(_c4567, 1)); + _out6 = vaddq_f32(_out6, vdupq_laneq_f32(_c4567, 2)); + _out7 = vaddq_f32(_out7, vdupq_laneq_f32(_c4567, 3)); + } + 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) + { + 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; + } + 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); + pp += 16; + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vadd_f32(_out0, vdup_lane_f32(vget_low_f32(_c0123), 0)); + _out1 = vadd_f32(_out1, vdup_lane_f32(vget_low_f32(_c0123), 1)); + _out2 = vadd_f32(_out2, vdup_lane_f32(vget_high_f32(_c0123), 0)); + _out3 = vadd_f32(_out3, vdup_lane_f32(vget_high_f32(_c0123), 1)); + _out4 = vadd_f32(_out4, vdup_lane_f32(vget_low_f32(_c4567), 0)); + _out5 = vadd_f32(_out5, vdup_lane_f32(vget_low_f32(_c4567), 1)); + _out6 = vadd_f32(_out6, vdup_lane_f32(vget_high_f32(_c4567), 0)); + _out7 = vadd_f32(_out7, vdup_lane_f32(vget_high_f32(_c4567), 1)); + } + 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) + { + 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; + } + 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]; + pp += 8; + 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++; + } + } +#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; + } + + float32x4_t _c0_broadcast = vdupq_n_f32(c0); + float32x4_t _c1_broadcast = vdupq_n_f32(c1); + float32x4_t _c2_broadcast = vdupq_n_f32(c2); + float32x4_t _c3_broadcast = vdupq_n_f32(c3); + float32x4_t _c01_broadcast = vcombine_f32(vdup_n_f32(c0), vdup_n_f32(c1)); + float32x4_t _c23_broadcast = vcombine_f32(vdup_n_f32(c2), vdup_n_f32(c3)); + float32x4_t _c0123_broadcast = vsetq_lane_f32(c1, _c0_broadcast, 1); + _c0123_broadcast = vsetq_lane_f32(c2, _c0123_broadcast, 2); + _c0123_broadcast = vsetq_lane_f32(c3, _c0123_broadcast, 3); + + 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); + pp += 32; + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0_broadcast); + _out1 = vaddq_f32(_out1, _c0_broadcast); + _out2 = vaddq_f32(_out2, _c0_broadcast); + _out3 = vaddq_f32(_out3, _c0_broadcast); + _out4 = vaddq_f32(_out4, _c0_broadcast); + _out5 = vaddq_f32(_out5, _c0_broadcast); + _out6 = vaddq_f32(_out6, _c0_broadcast); + _out7 = vaddq_f32(_out7, _c0_broadcast); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0_broadcast); + _out1 = vaddq_f32(_out1, _c0_broadcast); + _out2 = vaddq_f32(_out2, _c1_broadcast); + _out3 = vaddq_f32(_out3, _c1_broadcast); + _out4 = vaddq_f32(_out4, _c2_broadcast); + _out5 = vaddq_f32(_out5, _c2_broadcast); + _out6 = vaddq_f32(_out6, _c3_broadcast); + _out7 = vaddq_f32(_out7, _c3_broadcast); + } + 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))); + pC += 8; + } + if (broadcast_type_C == 4) + { + 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); + 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; + } + } + + 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; + } +#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); + pp += 16; + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0_broadcast); + _out1 = vaddq_f32(_out1, _c0_broadcast); + _out2 = vaddq_f32(_out2, _c0_broadcast); + _out3 = vaddq_f32(_out3, _c0_broadcast); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0_broadcast); + _out1 = vaddq_f32(_out1, _c1_broadcast); + _out2 = vaddq_f32(_out2, _c2_broadcast); + _out3 = vaddq_f32(_out3, _c3_broadcast); + } + 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))); + pC += 4; + } + if (broadcast_type_C == 4) + { + 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; + } + } + + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0_broadcast); + _out1 = vaddq_f32(_out1, _c0_broadcast); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c01_broadcast); + _out1 = vaddq_f32(_out1, _c23_broadcast); + } + if (broadcast_type_C == 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); + float32x4_t _cc0 = vcombine_f32(_c, _c); + _out0 = vaddq_f32(_out0, _cc0); + _out1 = vaddq_f32(_out1, _cc0); + pC += 2; + } + } + + 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; + } + for (; jj < max_jj; jj += 1) + { + float32x4_t _out0 = vld1q_f32(pp); + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0_broadcast); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0123_broadcast); + } + 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); + pC++; + } + if (broadcast_type_C == 4) + { + float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + _out0 = vaddq_f32(_out0, _cc0); + pC++; + } + } + + 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++; + } + } +#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 = vdupq_n_f32(0.f); + float32x4_t _c1 = vdupq_n_f32(0.f); +#endif + float c0 = 0.f; + float c1 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + 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; + 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; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; +#if __ARM_NEON +#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); + pp += 16; + + 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))); + pC += 8; + } + if (broadcast_type_C == 4) + { + 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); + 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; + } + } + + 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; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + pp += 8; + + 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))); + pC += 4; + } + if (broadcast_type_C == 4) + { + 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; + } + } + + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + float32x2_t _out1 = vld1_f32(pp + 2); + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0) + { + 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) + { + 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) + { + 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; + } + } + + 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; + } + for (; jj < max_jj; jj += 1) + { + float32x2_t _out0 = vld1_f32(pp); + pp += 2; + + 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); + pC++; + } + if (broadcast_type_C == 4) + { + _out0 = vadd_f32(_out0, vdup_n_f32(beta == 1.f ? pC[0] : pC[0] * beta)); + pC++; + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + } + + vst1_lane_f32(outptr0, _out0, 0); + vst1_lane_f32(outptr1, _out0, 1); + + outptr0++; + outptr1++; + } +#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]; + pp += 4; + + 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; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + float out10 = pp[1]; + pp += 2; + + 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++; + } + } + for (; ii < max_ii; ii++) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; +#if __ARM_NEON + float32x4_t _c0 = vdupq_n_f32(0.f); +#endif + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + 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; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + 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); + pp += 16; + + 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; + } + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + pp += 8; + + 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))); + pC += 8; + } + if (broadcast_type_C == 4) + { + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + 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; + } + } + + 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; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + pp += 4; + + 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))); + pC += 4; + } + if (broadcast_type_C == 4) + { + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + pC += 4; + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + } + + vst1q_f32(outptr0, _out0); + + outptr0 += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + pp += 2; + + 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) + { + 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) + { + float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c0); + pC += 2; + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + } + + vst1_f32(outptr0, _out0); + + outptr0 += 2; + } +#endif // __ARM_NEON + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + pp += 2; + + 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; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + pp++; + + 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; + pC++; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + } + + outptr0[0] = out00; + + outptr0++; + } + } +} + +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, 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)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; + 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 + 7 < max_jj; jj += 8) + { + 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); + + float32x4_t _out8 = vld1q_f32(pp + 32); + float32x4_t _out9 = vld1q_f32(pp + 36); + float32x4_t _outa = vld1q_f32(pp + 40); + float32x4_t _outb = vld1q_f32(pp + 44); + float32x4_t _outc = vld1q_f32(pp + 48); + float32x4_t _outd = vld1q_f32(pp + 52); + float32x4_t _oute = vld1q_f32(pp + 56); + float32x4_t _outf = vld1q_f32(pp + 60); + pp += 64; + + if (pC) + { + if (broadcast_type_C <= 2) + { + float32x4_t _c0123; + float32x4_t _c4567; + if (broadcast_type_C == 0) + { + _c0123 = vdupq_n_f32(c0); + _c4567 = _c0123; + } + else + { + _c0123 = vmulq_n_f32(vld1q_f32(pC), beta); + _c4567 = vmulq_n_f32(vld1q_f32(pC + 4), beta); + } + _out0 = vaddq_f32(_out0, vdupq_laneq_f32(_c0123, 0)); + _out8 = vaddq_f32(_out8, vdupq_laneq_f32(_c0123, 0)); + _out1 = vaddq_f32(_out1, vdupq_laneq_f32(_c0123, 1)); + _out9 = vaddq_f32(_out9, vdupq_laneq_f32(_c0123, 1)); + _out2 = vaddq_f32(_out2, vdupq_laneq_f32(_c0123, 2)); + _outa = vaddq_f32(_outa, vdupq_laneq_f32(_c0123, 2)); + _out3 = vaddq_f32(_out3, vdupq_laneq_f32(_c0123, 3)); + _outb = vaddq_f32(_outb, vdupq_laneq_f32(_c0123, 3)); + _out4 = vaddq_f32(_out4, vdupq_laneq_f32(_c4567, 0)); + _outc = vaddq_f32(_outc, vdupq_laneq_f32(_c4567, 0)); + _out5 = vaddq_f32(_out5, vdupq_laneq_f32(_c4567, 1)); + _outd = vaddq_f32(_outd, vdupq_laneq_f32(_c4567, 1)); + _out6 = vaddq_f32(_out6, vdupq_laneq_f32(_c4567, 2)); + _oute = vaddq_f32(_oute, vdupq_laneq_f32(_c4567, 2)); + _out7 = vaddq_f32(_out7, vdupq_laneq_f32(_c4567, 3)); + _outf = vaddq_f32(_outf, vdupq_laneq_f32(_c4567, 3)); + } + if (broadcast_type_C == 3) + { + float32x4_t _c0 = vld1q_f32(pC); + float32x4_t _c1 = vld1q_f32(pC + c_hstep); + float32x4_t _c2 = vld1q_f32(pC + c_hstep * 2); + float32x4_t _c3 = vld1q_f32(pC + c_hstep * 3); + float32x4_t _c4 = vld1q_f32(pC + c_hstep * 4); + float32x4_t _c5 = vld1q_f32(pC + c_hstep * 5); + float32x4_t _c6 = vld1q_f32(pC + c_hstep * 6); + float32x4_t _c7 = vld1q_f32(pC + c_hstep * 7); + if (beta == 1.f) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c1); + _out2 = vaddq_f32(_out2, _c2); + _out3 = vaddq_f32(_out3, _c3); + _out4 = vaddq_f32(_out4, _c4); + _out5 = vaddq_f32(_out5, _c5); + _out6 = vaddq_f32(_out6, _c6); + _out7 = vaddq_f32(_out7, _c7); + } + else + { + _out0 = vmlaq_n_f32(_out0, _c0, beta); + _out1 = vmlaq_n_f32(_out1, _c1, beta); + _out2 = vmlaq_n_f32(_out2, _c2, beta); + _out3 = vmlaq_n_f32(_out3, _c3, beta); + _out4 = vmlaq_n_f32(_out4, _c4, beta); + _out5 = vmlaq_n_f32(_out5, _c5, beta); + _out6 = vmlaq_n_f32(_out6, _c6, beta); + _out7 = vmlaq_n_f32(_out7, _c7, beta); + } + _c0 = vld1q_f32(pC + 4); + _c1 = vld1q_f32(pC + c_hstep + 4); + _c2 = vld1q_f32(pC + c_hstep * 2 + 4); + _c3 = vld1q_f32(pC + c_hstep * 3 + 4); + _c4 = vld1q_f32(pC + c_hstep * 4 + 4); + _c5 = vld1q_f32(pC + c_hstep * 5 + 4); + _c6 = vld1q_f32(pC + c_hstep * 6 + 4); + _c7 = vld1q_f32(pC + c_hstep * 7 + 4); + if (beta == 1.f) + { + _out8 = vaddq_f32(_out8, _c0); + _out9 = vaddq_f32(_out9, _c1); + _outa = vaddq_f32(_outa, _c2); + _outb = vaddq_f32(_outb, _c3); + _outc = vaddq_f32(_outc, _c4); + _outd = vaddq_f32(_outd, _c5); + _oute = vaddq_f32(_oute, _c6); + _outf = vaddq_f32(_outf, _c7); + } + else + { + _out8 = vmlaq_n_f32(_out8, _c0, beta); + _out9 = vmlaq_n_f32(_out9, _c1, beta); + _outa = vmlaq_n_f32(_outa, _c2, beta); + _outb = vmlaq_n_f32(_outb, _c3, beta); + _outc = vmlaq_n_f32(_outc, _c4, beta); + _outd = vmlaq_n_f32(_outd, _c5, beta); + _oute = vmlaq_n_f32(_oute, _c6, beta); + _outf = vmlaq_n_f32(_outf, _c7, beta); + } + pC += 8; + } + if (broadcast_type_C == 4) + { + float32x4_t _c = vld1q_f32(pC); + if (beta != 1.f) + _c = vmulq_n_f32(_c, 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); + _c = vld1q_f32(pC + 4); + if (beta != 1.f) + _c = vmulq_n_f32(_c, beta); + _out8 = vaddq_f32(_out8, _c); + _out9 = vaddq_f32(_out9, _c); + _outa = vaddq_f32(_outa, _c); + _outb = vaddq_f32(_outb, _c); + _outc = vaddq_f32(_outc, _c); + _outd = vaddq_f32(_outd, _c); + _oute = vaddq_f32(_oute, _c); + _outf = vaddq_f32(_outf, _c); + pC += 8; + } + } + 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); + _out8 = vmulq_n_f32(_out8, alpha); + _out9 = vmulq_n_f32(_out9, alpha); + _outa = vmulq_n_f32(_outa, alpha); + _outb = vmulq_n_f32(_outb, alpha); + _outc = vmulq_n_f32(_outc, alpha); + _outd = vmulq_n_f32(_outd, alpha); + _oute = vmulq_n_f32(_oute, alpha); + _outf = vmulq_n_f32(_outf, 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); + + transpose4x4_ps(_out8, _out9, _outa, _outb); + transpose4x4_ps(_outc, _outd, _oute, _outf); + vst1q_f32(outptr4, _out8); + vst1q_f32(outptr4 + 4, _outc); + vst1q_f32(outptr5, _out9); + vst1q_f32(outptr5 + 4, _outd); + vst1q_f32(outptr6, _outa); + vst1q_f32(outptr6 + 4, _oute); + vst1q_f32(outptr7, _outb); + vst1q_f32(outptr7 + 4, _outf); + outptr0 += out_hstep * 8; + outptr1 += out_hstep * 8; + outptr2 += out_hstep * 8; + outptr3 += out_hstep * 8; + outptr4 += out_hstep * 8; + outptr5 += out_hstep * 8; + outptr6 += out_hstep * 8; + outptr7 += out_hstep * 8; + } + float32x4_t _c0123 = vdupq_n_f32(c0); + _c0123 = vsetq_lane_f32(c1, _c0123, 1); + _c0123 = vsetq_lane_f32(c2, _c0123, 2); + _c0123 = vsetq_lane_f32(c3, _c0123, 3); + float32x4_t _c4567 = vdupq_n_f32(c4); + _c4567 = vsetq_lane_f32(c5, _c4567, 1); + _c4567 = vsetq_lane_f32(c6, _c4567, 2); + _c4567 = vsetq_lane_f32(c7, _c4567, 3); + 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); + pp += 32; + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, vdupq_laneq_f32(_c0123, 0)); + _out1 = vaddq_f32(_out1, vdupq_laneq_f32(_c0123, 1)); + _out2 = vaddq_f32(_out2, vdupq_laneq_f32(_c0123, 2)); + _out3 = vaddq_f32(_out3, vdupq_laneq_f32(_c0123, 3)); + _out4 = vaddq_f32(_out4, vdupq_laneq_f32(_c4567, 0)); + _out5 = vaddq_f32(_out5, vdupq_laneq_f32(_c4567, 1)); + _out6 = vaddq_f32(_out6, vdupq_laneq_f32(_c4567, 2)); + _out7 = vaddq_f32(_out7, vdupq_laneq_f32(_c4567, 3)); + } + 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) + { + 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; + } + 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); + pp += 16; + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vadd_f32(_out0, vdup_lane_f32(vget_low_f32(_c0123), 0)); + _out1 = vadd_f32(_out1, vdup_lane_f32(vget_low_f32(_c0123), 1)); + _out2 = vadd_f32(_out2, vdup_lane_f32(vget_high_f32(_c0123), 0)); + _out3 = vadd_f32(_out3, vdup_lane_f32(vget_high_f32(_c0123), 1)); + _out4 = vadd_f32(_out4, vdup_lane_f32(vget_low_f32(_c4567), 0)); + _out5 = vadd_f32(_out5, vdup_lane_f32(vget_low_f32(_c4567), 1)); + _out6 = vadd_f32(_out6, vdup_lane_f32(vget_high_f32(_c4567), 0)); + _out7 = vadd_f32(_out7, vdup_lane_f32(vget_high_f32(_c4567), 1)); + } + 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) + { + 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; + } + for (; jj < max_jj; jj++) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + pp += 8; + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, _c0123); + _out1 = vaddq_f32(_out1, _c4567); + } + if (broadcast_type_C == 3) + { + float32x4_t _c0 = vdupq_n_f32(pC[0]); + _c0 = vsetq_lane_f32(pC[c_hstep], _c0, 1); + _c0 = vsetq_lane_f32(pC[c_hstep * 2], _c0, 2); + _c0 = vsetq_lane_f32(pC[c_hstep * 3], _c0, 3); + float32x4_t _c1 = vdupq_n_f32(pC[c_hstep * 4]); + _c1 = vsetq_lane_f32(pC[c_hstep * 5], _c1, 1); + _c1 = vsetq_lane_f32(pC[c_hstep * 6], _c1, 2); + _c1 = vsetq_lane_f32(pC[c_hstep * 7], _c1, 3); + _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) + { + 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; + } + } +#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); + pp += 32; + + 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))); + pC += 8; + } + if (broadcast_type_C == 4) + { + 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); + 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; + } + } + + 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) + { + 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); + _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; + } +#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); + pp += 16; + + 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))); + pC += 4; + } + if (broadcast_type_C == 4) + { + 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; + } + } + + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 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); + float32x4_t _cc0 = vcombine_f32(_c, _c); + _out0 = vaddq_f32(_out0, _cc0); + _out1 = vaddq_f32(_out1, _cc0); + pC += 2; + } + } + + 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; + } + for (; jj < max_jj; jj += 1) + { + float32x4_t _out0 = vld1q_f32(pp); + pp += 4; + + 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); + pC++; + } + if (broadcast_type_C == 4) + { + float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + _out0 = vaddq_f32(_out0, _cc0); + pC++; + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + } + + vst1q_f32(outptr0, _out0); + + outptr0 += out_hstep; + } + } +#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) + { + 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) + { + c0 = pC[i + ii]; + if (beta != 1.f) + 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) + pC = (const float*)C + j; + } + + int jj = 0; +#if __ARM_NEON +#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); + pp += 16; + + 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))); + pC += 8; + } + if (broadcast_type_C == 4) + { + 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); + 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; + } + } + + 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) + { + 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); + _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; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + pp += 8; + + 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))); + pC += 4; + } + if (broadcast_type_C == 4) + { + 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; + } + } + + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + float32x2_t _out1 = vld1_f32(pp + 2); + pp += 4; + + if (pC) + { + if (broadcast_type_C == 3) + { + 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) + { + 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; + } + } + + 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; + } + for (; jj < max_jj; jj += 1) + { + float32x2_t _out0 = vld1_f32(pp); + pp += 2; + + 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); + pC++; + } + if (broadcast_type_C == 4) + { + _out0 = vadd_f32(_out0, vdup_n_f32(beta == 1.f ? pC[0] : pC[0] * beta)); + pC++; + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + } + + vst1_f32(outptr0, _out0); + + outptr0 += out_hstep; + } +#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]; + pp += 4; + + 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; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + float out10 = pp[1]; + pp += 2; + + 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; + } + } + for (; ii < max_ii; ii++) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; +#if __ARM_NEON + float32x4_t _c0 = vdupq_n_f32(0.f); +#endif + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + 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; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + 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); + pp += 16; + + 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; + } + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + pp += 8; + + 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))); + pC += 8; + } + if (broadcast_type_C == 4) + { + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + 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; + } + } + + 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; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + pp += 4; + + 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))); + pC += 4; + } + if (broadcast_type_C == 4) + { + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + pC += 4; + } + } + + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + pp += 2; + + 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) + { + 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) + { + float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c0); + pC += 2; + } + } + + 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; + } +#endif // __ARM_NEON + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + pp += 2; + + 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[out_hstep] = out01; + + outptr0 += out_hstep * 2; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + pp++; + + 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; + pC++; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + } + + outptr0[0] = out00; + + outptr0 += out_hstep; + } + } +} diff --git a/src/layer/arm/multiheadattention_arm.cpp b/src/layer/arm/multiheadattention_arm.cpp index a3bbaf98785..d23123522e6 100644 --- a/src/layer/arm/multiheadattention_arm.cpp +++ b/src/layer/arm/multiheadattention_arm.cpp @@ -32,10 +32,367 @@ 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; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -53,18 +410,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 +462,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 +511,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) + { + destroy_pipeline(opt); + return ret; + } + ret = k_gemm->create_pipeline(opt); + if (ret != 0) { - k_weight_data.release(); - k_bias_data.release(); + 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 +560,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 +607,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 +666,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 +714,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; @@ -269,10 +756,20 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) int MultiHeadAttention_arm::destroy_pipeline(const Option& _opt) { if (weight_block_quantize) - return 0; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 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 +779,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 +799,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; } @@ -337,7 +845,17 @@ 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) - return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + } int q_blob_i = 0; int k_blob_i = 0; @@ -355,7 +873,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 +883,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 +942,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 +952,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 +980,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 +1031,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 +1059,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 +1105,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 c99988b0319..b70c9ceab55 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 2023bcf7011..5dbdf6c79e2 100644 --- a/src/layer/gemm.cpp +++ b/src/layer/gemm.cpp @@ -8,51 +8,172 @@ namespace ncnn { -static bool gemm_is_weight_block_quantize(int quantize_term) +int Gemm::get_weight_block_quantize_params(int& weight_bits, int& block_size, bool& has_input_scale) const { - const int weight_bits = quantize_term / 100; + weight_bits = quantize_term / 100; const int format_code = quantize_term % 100 / 10; const int block_size_code = quantize_term % 10; if (weight_bits != 4 && weight_bits != 6 && weight_bits != 8) - return false; + return -1; if (format_code != 0 && format_code != 1) - return false; + return -1; if (block_size_code < 0 || block_size_code > 2) - return false; + return -1; - return true; + block_size = block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + has_input_scale = format_code == 1; + + return 0; } #if NCNN_WEIGHT_QUANT -static bool gemm_weight_quantize_has_input_scale(int quantize_term) +static int gemm_weight_quantize_packed_k_bytes(int constantK, int weight_bits) { - return quantize_term % 100 / 10 == 1; + if (constantK <= 0 || weight_bits <= 0) + return -1; + + const size_t packed_k_bytes = ((size_t)constantK * weight_bits + 7) / 8; + if (packed_k_bytes > (size_t)INT_MAX) + return -1; + + return (int)packed_k_bytes; } -static int gemm_weight_quantize_bits(int quantize_term) +static inline signed char weight_block_quantize_float2int8(float v) { - return quantize_term / 100; + int int32 = static_cast(round(v)); + if (int32 > 127) return 127; + if (int32 < -127) return -127; + return (signed char)int32; } -static int gemm_weight_quantize_block_size(int quantize_term) +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_size_code = quantize_term % 10; - return block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + 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; + } + + const float scale = 127.f / absmax; + 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]; + outptr[k] = weight_block_quantize_float2int8(v * scale); + } + } } -static int gemm_weight_quantize_packed_k_bytes(int constantK, int weight_bits) +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) { - if (constantK <= 0 || weight_bits <= 0) - return -1; + const int block_count = (K + block_size - 1) / block_size; - const size_t packed_k_bytes = ((size_t)constantK * weight_bits + 7) / 8; - if (packed_k_bytes > (size_t)INT_MAX) - return -1; + Mat A_int8; + A_int8.create(K, M, (size_t)1u, opt.workspace_allocator); + if (A_int8.empty()) + return -100; - return (int)packed_k_bytes; + 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; + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + + #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) + ((float*)top_blob)[(size_t)j * out_hstep + output_m_offset + i] = sum; + else + ((float*)top_blob)[(size_t)(output_m_offset + i) * out_hstep + j] = sum; + } + + return 0; } #endif // NCNN_WEIGHT_QUANT @@ -82,7 +203,10 @@ int Gemm::load_param(const ParamDict& pd) output_elemtype = pd.get(13, 0); output_transpose = pd.get(14, 0); quantize_term = pd.get(18, 0); - weight_block_quantize = gemm_is_weight_block_quantize(quantize_term); + int weight_bits; + int block_size; + bool has_input_scale; + weight_block_quantize = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) == 0; constant_TILE_M = pd.get(20, 0); constant_TILE_N = pd.get(21, 0); constant_TILE_K = pd.get(22, 0); @@ -102,13 +226,13 @@ int Gemm::load_param(const ParamDict& pd) if (weight_block_quantize) { #if NCNN_WEIGHT_QUANT - if (constantA != 0 || constantB != 1 || transA != 0 || transB != 1) + 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 && weight_bits != 8) || 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; @@ -164,6 +288,14 @@ int Gemm::load_param(const ParamDict& pd) int Gemm::load_model(const ModelBin& mb) { +#if NCNN_WEIGHT_QUANT + int weight_bits = 0; + int block_size = 0; + bool has_input_scale = false; + if (weight_block_quantize && get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; +#endif + if (constantA == 1) { if (transA == 0) @@ -181,7 +313,6 @@ int Gemm::load_model(const ModelBin& mb) #if NCNN_WEIGHT_QUANT else if (weight_block_quantize) { - const int weight_bits = gemm_weight_quantize_bits(quantize_term); const int packed_k_bytes = gemm_weight_quantize_packed_k_bytes(constantK, weight_bits); if (packed_k_bytes < 0) return -100; @@ -213,14 +344,13 @@ int Gemm::load_model(const ModelBin& mb) #if NCNN_WEIGHT_QUANT if (weight_block_quantize) { - const int block_size = gemm_weight_quantize_block_size(quantize_term); const int block_count = (constantK + block_size - 1) / block_size; B_data_quantize_scales = mb.load(block_count, constantN, 1); if (B_data_quantize_scales.empty()) return -100; - if (gemm_weight_quantize_has_input_scale(quantize_term)) + if (has_input_scale) { B_data_input_scales = mb.load(constantK, 1); if (B_data_input_scales.empty()) @@ -344,21 +474,29 @@ 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"); return -1; } - const int weight_bits = gemm_weight_quantize_bits(quantize_term); - const int block_size = gemm_weight_quantize_block_size(quantize_term); + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; const int packed_k_bytes = gemm_weight_quantize_packed_k_bytes(constantK, weight_bits); - 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 +562,31 @@ 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) + { + if (output_N1M) + top_blob.create(M, 1, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + } + else + { + if (output_N1M) + top_blob.create(N, 1, M, (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/gemm.h b/src/layer/gemm.h index af0bbc2f025..e36698fb028 100644 --- a/src/layer/gemm.h +++ b/src/layer/gemm.h @@ -22,6 +22,8 @@ class Gemm : public Layer virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; protected: + int get_weight_block_quantize_params(int& weight_bits, int& block_size, bool& has_input_scale) const; + #if NCNN_WEIGHT_QUANT int forward_weight_block_quantize(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; #endif diff --git a/src/layer/loongarch/gemm_loongarch.cpp b/src/layer/loongarch/gemm_loongarch.cpp index 72832080bcf..dcbdf345c0d 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,18 +7369,160 @@ static int gemm_AT_BT_loongarch(const Mat& AT, const Mat& BT, const Mat& C, Mat& return 0; } -int Gemm_loongarch::create_pipeline(const Option& opt) +#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) { - AT_data.release(); - BT_data.release(); - CT_data.release(); - nT = 0; + 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, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + 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); + 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; + + const int nn_MK = nn_M * nn_K; + #pragma omp parallel for num_threads(nT) + 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; + + 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(local_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, 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); + } + + #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 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); + 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 local_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(local_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, 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); + 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, 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_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()); + + 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); + 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; + Mat AT_tile(max_kk, max_ii, (signed char*)AT_channel + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile(local_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, 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); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + } + } + } + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + +int Gemm_loongarch::create_pipeline(const Option& opt) +{ if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + int weight_bits; + int block_size; + bool has_input_scale; + int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits == 8) + return create_pipeline_wq_int8(opt); +#endif return 0; } + AT_data.release(); + BT_data.release(); + CT_data.release(); + nT = 0; + #if NCNN_INT8 if (quantize_term) { @@ -7525,10 +7671,160 @@ int Gemm_loongarch::create_pipeline(const Option& opt) return 0; } +int Gemm_loongarch::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_loongarch::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; + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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); + if (ret != 0) + return ret; + if (BT_data_packed.empty() || BT_data_packed_descales.empty()) + return -100; + + 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; + } + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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) + { + if (output_N1M) + top_blob.create(M, 1, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + } + else + { + if (output_N1M) + top_blob.create(N, 1, M, (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_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 + int weight_bits; + int block_size; + bool has_input_scale; + int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits == 8) + 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 050557e5ba4..3435e635746 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 00000000000..dcf20f1b8e7 --- /dev/null +++ b/src/layer/loongarch/gemm_wq_int8.h @@ -0,0 +1,7327 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +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, const Option& opt) +{ + const int block_count = (K + block_size - 1) / block_size; + Mat BT_packed; + BT_packed.create(N * K, (size_t)1u, opt.blob_allocator); + if (BT_packed.empty()) + return -100; + + Mat BT_packed_descales; + BT_packed_descales.create(N * block_count, (size_t)4u, opt.blob_allocator); + if (BT_packed_descales.empty()) + return -100; + BT_packed.cstep = (size_t)N * K; + BT_packed_descales.cstep = (size_t)N * block_count; + + int j = 0; +#if __loongarch_sx +#if __loongarch_asx + const int nn8 = N / 8; + const int j8 = j; + j += nn8 * 8; +#endif // __loongarch_asx + const int nn4 = (N - j) / 4; + const int j4 = j; + j += nn4 * 4; +#endif // __loongarch_sx + const int nn2 = (N - j) / 2; + const int j2 = j; + j += nn2 * 2; + const int nn1 = N - j; + const int j1 = j; + + #pragma omp parallel num_threads(opt.num_threads) + { +#if __loongarch_sx +#if __loongarch_asx + #pragma omp for + for (int p = 0; p < nn8; p++) + { + const int j = j8 + p * 8; + signed char* pp = (signed char*)BT_packed + (size_t)j * K; + float* pd = (float*)BT_packed_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 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++; + } + } +#endif // __loongarch_asx + #pragma omp for + for (int p = 0; p < nn4; p++) + { + const int j = j4 + p * 4; + signed char* pp = (signed char*)BT_packed + (size_t)j * K; + float* pd = (float*)BT_packed_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) + { + __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) + { + 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++; + } + } +#endif // __loongarch_sx + #pragma omp for + for (int p = 0; p < nn2; p++) + { + const int j = j2 + p * 2; + signed char* pp = (signed char*)BT_packed + (size_t)j * K; + float* pd = (float*)BT_packed_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) + { + 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) + { + 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; + } + + *pd++ = 1.f / *s0++; + *pd++ = 1.f / *s1++; + } + } + + #pragma omp for + for (int p = 0; p < nn1; p++) + { + const int j = j1 + p; + signed char* pp = (signed char*)BT_packed + (size_t)j * K; + float* pd = (float*)BT_packed_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); + 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; + 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++; + + *pd++ = 1.f / *s0++; + } + } + } + + 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 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 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 __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + 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; + 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = p0 + kk0; + const float* p1g = p1 + kk0; + const float* p2g = p2 + kk0; + const float* p3g = p3 + kk0; + const float* p4g = p4 + kk0; + const float* p5g = p5 + kk0; + const float* p6g = p6 + kk0; + const float* p7g = p7 + kk0; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 0; + 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); + 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_kk0; kk += 4) + { + __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); + _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)); + 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); + 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_kk0; kk++) + { + 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; + 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; + + 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_kk0; kk += 4) + { + __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); + _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; + p0q += 4; + p1q += 4; + p2q += 4; + p3q += 4; + p4q += 4; + p5q += 4; + p6q += 4; + p7q += 4; + if (psq) + psq += 4; + } + for (; kk < max_kk0; kk++) + { + 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; + } + } + } + for (; ii + 3 < max_ii; ii += 4) + { + 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; + 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = p0 + kk0; + const float* p1g = p1 + kk0; + const float* p2g = p2 + kk0; + const float* p3g = p3 + kk0; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 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_kk0; kk += 4) + { + __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) + { + __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); + _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)); + 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); + float absmax2 = __lsx_reduce_fmax_s(_absmax2); + float absmax3 = __lsx_reduce_fmax_s(_absmax3); + for (; kk < max_kk0; kk++) + { + 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; + 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_kk0; kk += 4) + { + __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) + { + __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); + _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; + p0q += 4; + p1q += 4; + p2q += 4; + p3q += 4; + if (psq) + psq += 4; + } + if (kk + 1 < max_kk0) + { + 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_kk0) + { + 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 = 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; + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + + for (int g = 0; g < block_count; g++) + { + const int kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = p0 + kk0; + const float* p1g = p1 + kk0; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 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_kk0; kk += 4) + { + __m128 _v0 = (__m128)__lsx_vld(p0a, 0); + __m128 _v1 = (__m128)__lsx_vld(p1a, 0); + if (psa) + { + __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_kk0; 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; + __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_kk0; kk += 4) + { + __m128 _v0 = (__m128)__lsx_vld(p0q, 0); + __m128 _v1 = (__m128)__lsx_vld(p1q, 0); + if (psq) + { + __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_kk0) + { + 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_kk0) + { + 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = p0 + kk0; + const float* p1g = p1 + kk0; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 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_kk0; 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_kk0; 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_kk0) + { + 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_kk0) + { + 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* 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = p0 + kk0; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 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); +#if __loongarch_asx + const __m256i _abs_mask256 = __lasx_xvreplgr2vr_w(0x7fffffff); + __m256 _absmax256 = (__m256)__lasx_xvldi(0); + for (; kk + 7 < max_kk0; kk += 8) + { + __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 // __loongarch_asx + __m128 _absmax128 = __lsx_vreplfr2vr_s(absmax); + for (; kk + 3 < max_kk0; kk += 4) + { + __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 // __loongarch_sx + for (; kk < max_kk0; kk++) + { + float v = *p0a++; + if (psa) + v *= *psa++; + absmax = std::max(absmax, fabsf(v)); + } + + if (absmax == 0.f) + { + *pd++ = 0.f; + for (int kk = 0; kk < max_kk0; kk++) + *pp++ = 0; + continue; + } + + 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 + __m256 _scale256 = (__m256)__lasx_xvreplfr2vr_s(scale); + for (; kk + 7 < max_kk0; kk += 8) + { + __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)), pp, 0, 0); + pp += 8; + p0q += 8; + if (psq) + psq += 8; + } +#endif // __loongarch_asx + __m128 _scale128 = __lsx_vreplfr2vr_s(scale); + for (; kk + 3 < max_kk0; kk += 4) + { + __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), pp, 0, 0); + pp += 4; + p0q += 4; + if (psq) + psq += 4; + } +#endif // __loongarch_sx + for (; kk < max_kk0; kk++) + { + 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 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 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 __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = ptrA + (size_t)kk0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 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_kk0; kk++) + { + __m128 _v0 = (__m128)__lsx_vld(p0a, 0); + __m128 _v1 = (__m128)__lsx_vld(p0a + 4, 0); + if (psa) + { + __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]; + __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; + + 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}; + __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_kk0; kk += 4) + { + const float* p0 = p0q; + 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 (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); + *((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; + p0q += (size_t)4 * A_hstep; + if (psq) + psq += 4; + } + for (; kk < max_kk0; kk++) + { + __m128 _p0 = (__m128)__lsx_vld(p0q, 0); + __m128 _p1 = (__m128)__lsx_vld(p0q + 4, 0); + if (psq) + { + __m128 _s = __lsx_vreplfr2vr_s(*psq++); + _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; + p0q += A_hstep; + } + } + } + for (; ii + 3 < max_ii; ii += 4) + { + 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = ptrA + (size_t)kk0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 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_kk0; kk++) + { + __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]; + __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; + + const float scales[4] = { + 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] + }; + __m128 _scale = (__m128)__lsx_vld(scales, 0); + + const float* p0q = p0g; + const float* psq = sg; + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + const float* p0 = p0q; + 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 (psq) + { + _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_kk0) + { + __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(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); + 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; + p0q += (size_t)2 * A_hstep; + if (psq) + psq += 2; + kk += 2; + } + if (kk < max_kk0) + { + __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); + pp[2] = (signed char)(q >> 16); + pp[3] = (signed char)(q >> 24); + pp += 4; + } + } + } + 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; + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + + for (int g = 0; g < block_count; g++) + { + const int kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = ptrA + (size_t)kk0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 0; + __m128 _absmax = (__m128)__lsx_vldi(0); + const float* p0a = p0g; + const float* psa = sg; + for (int kk = 0; kk < max_kk0; kk++) + { + __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; + 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_kk0; 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_kk0) + { + 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_kk0) + { + 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = ptrA + (size_t)kk0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 0; + float absmax0 = 0.f; + float absmax1 = 0.f; + const float* p0a = p0g; + const float* psa = sg; + for (int kk = 0; kk < max_kk0; 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_kk0; 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_kk0) + { + 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_kk0) + { + 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* 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 kk0 = g * block_size; + const int max_kk0 = std::min(max_kk - kk0, block_size); + const float* p0g = p0 + (size_t)kk0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + kk0 : 0; + + float absmax = 0.f; + const float* p0a = p0g; + const float* psa = sg; + for (int kk = 0; kk < max_kk0; kk++) + { + float v = *p0a; + if (psa) + v *= *psa++; + absmax = std::max(absmax, fabsf(v)); + p0a += A_hstep; + } + + if (absmax == 0.f) + { + *pd++ = 0.f; + for (int kk = 0; kk < max_kk0; kk++) + *pp++ = 0; + continue; + } + + const float scale = 127.f / absmax; + *pd++ = absmax / 127.f; + + const float* p0q = p0g; + const float* psq = sg; + for (int kk = 0; kk < max_kk0; kk++) + { + 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 max_kk, int K, int block_size) +{ + const signed char* pAT = AT_tile; + const int A_hstep = max_kk; + const float* pAT_descales = AT_descales_tile; + const int A_descales_hstep = (max_kk + block_size - 1) / block_size; + const signed char* pBT = BT_tile; + const float* pBT_descales = BT_descales_tile; + float* outptr = topT_tile; + const int block_count = (K + block_size - 1) / block_size; + const int block_start = k / block_size; + + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __loongarch_asx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)8 * k; + const float* pB_descales = pB_descales_panel + (size_t)8 * block_start; + __m256 _out0; + __m256 _out1; + __m256 _out2; + __m256 _out3; + __m256 _out4; + __m256 _out5; + __m256 _out6; + __m256 _out7; + if (k == 0) + { + _out0 = (__m256)__lasx_xvldi(0); + _out1 = (__m256)__lasx_xvldi(0); + _out2 = (__m256)__lasx_xvldi(0); + _out3 = (__m256)__lasx_xvldi(0); + _out4 = (__m256)__lasx_xvldi(0); + _out5 = (__m256)__lasx_xvldi(0); + _out6 = (__m256)__lasx_xvldi(0); + _out7 = (__m256)__lasx_xvldi(0); + } + else + { + _out0 = (__m256)__lasx_xvld(outptr, 0); + _out1 = (__m256)__lasx_xvld(outptr + 8, 0); + _out2 = (__m256)__lasx_xvld(outptr + 16, 0); + _out3 = (__m256)__lasx_xvld(outptr + 24, 0); + _out4 = (__m256)__lasx_xvld(outptr + 32, 0); + _out5 = (__m256)__lasx_xvld(outptr + 40, 0); + _out6 = (__m256)__lasx_xvld(outptr + 48, 0); + _out7 = (__m256)__lasx_xvld(outptr + 56, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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); + _s0 = __lasx_xvmaddwod_h_b(_s0, _pA, _pB0); + _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB0); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s1, _s1)); + _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; + } + if (kk + 1 < max_kk0) + { + __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); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(_s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(_s1)); + _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; + } + if (kk < max_kk0) + { + __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); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(_s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(_s1)); + _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; + } + + __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)); + _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); + 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_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __loongarch_asx + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + __m128 _out00; + __m128 _out01; + __m128 _out10; + __m128 _out11; + __m128 _out20; + __m128 _out21; + __m128 _out30; + __m128 _out31; + if (k == 0) + { + _out00 = (__m128)__lsx_vldi(0); + _out01 = (__m128)__lsx_vldi(0); + _out10 = (__m128)__lsx_vldi(0); + _out11 = (__m128)__lsx_vldi(0); + _out20 = (__m128)__lsx_vldi(0); + _out21 = (__m128)__lsx_vldi(0); + _out30 = (__m128)__lsx_vldi(0); + _out31 = (__m128)__lsx_vldi(0); + } + else + { + _out00 = (__m128)__lsx_vld(outptr, 0); + _out01 = (__m128)__lsx_vld(outptr + 4, 0); + _out10 = (__m128)__lsx_vld(outptr + 8, 0); + _out11 = (__m128)__lsx_vld(outptr + 12, 0); + _out20 = (__m128)__lsx_vld(outptr + 16, 0); + _out21 = (__m128)__lsx_vld(outptr + 20, 0); + _out30 = (__m128)__lsx_vld(outptr + 24, 0); + _out31 = (__m128)__lsx_vld(outptr + 28, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __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_kk0) + { + __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)); + _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); + 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_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + __m128 _out00; + __m128 _out01; + __m128 _out10; + __m128 _out11; + if (k == 0) + { + _out00 = (__m128)__lsx_vldi(0); + _out01 = (__m128)__lsx_vldi(0); + _out10 = (__m128)__lsx_vldi(0); + _out11 = (__m128)__lsx_vldi(0); + } + else + { + _out00 = (__m128)__lsx_vld(outptr, 0); + _out01 = (__m128)__lsx_vld(outptr + 4, 0); + _out10 = (__m128)__lsx_vld(outptr + 8, 0); + _out11 = (__m128)__lsx_vld(outptr + 12, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __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_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); + __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_kk0) + { + __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_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)); + _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; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + __m128 _out0; + __m128 _out1; + if (k == 0) + { + _out0 = (__m128)__lsx_vldi(0); + _out1 = (__m128)__lsx_vldi(0); + } + else + { + _out0 = (__m128)__lsx_vld(outptr, 0); + _out1 = (__m128)__lsx_vld(outptr + 4, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + __m128i _sum1 = __lsx_vreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA0 = __lsx_vld(pA, 0); + __m128i _pA1 = __lsx_vld(pA + 16, 0); + __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)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s1, _s1)); + pA += 32; + pB += 4; + } + if (kk + 1 < max_kk0) + { + __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_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)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pA += 16; + pB += 2; + kk += 2; + } + if (kk < max_kk0) + { + __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; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep * 8; + pAT_descales += A_descales_hstep * 8; + } + for (; ii + 3 < max_ii; ii += 4) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __loongarch_asx + for (; jj + 15 < max_jj; jj += 16) + { + const signed char* pB0 = pB_panel + (size_t)8 * k; + const signed char* pB1 = pB_panel + (size_t)8 * K + (size_t)8 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)8 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)8 * block_count + (size_t)8 * block_start; + __m256 _out00; + __m256 _out01; + __m256 _out10; + __m256 _out11; + __m256 _out20; + __m256 _out21; + __m256 _out30; + __m256 _out31; + if (k == 0) + { + _out00 = (__m256)__lasx_xvldi(0); + _out01 = (__m256)__lasx_xvldi(0); + _out10 = (__m256)__lasx_xvldi(0); + _out11 = (__m256)__lasx_xvldi(0); + _out20 = (__m256)__lasx_xvldi(0); + _out21 = (__m256)__lasx_xvldi(0); + _out30 = (__m256)__lasx_xvldi(0); + _out31 = (__m256)__lasx_xvldi(0); + } + else + { + _out00 = (__m256)__lasx_xvld(outptr, 0); + _out01 = (__m256)__lasx_xvld(outptr + 8, 0); + _out10 = (__m256)__lasx_xvld(outptr + 16, 0); + _out11 = (__m256)__lasx_xvld(outptr + 24, 0); + _out20 = (__m256)__lasx_xvld(outptr + 32, 0); + _out21 = (__m256)__lasx_xvld(outptr + 40, 0); + _out30 = (__m256)__lasx_xvld(outptr + 48, 0); + _out31 = (__m256)__lasx_xvld(outptr + 56, 0); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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); + __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(_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, _pB0), _pA1, _pB0); + _sum20 = __lasx_xvadd_w(_sum20, __lasx_xvhaddw_w_h(_s, _s)); + _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(_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_kk0) + { + __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_kk0) + { + __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 _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(_pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA1, _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA1, _pB1); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA2, _pB0); + _sum20 = __lasx_xvadd_w(_sum20, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA2, _pB1); + _sum21 = __lasx_xvadd_w(_sum21, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA3, _pB0); + _sum30 = __lasx_xvadd_w(_sum30, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA3, _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; + } + + __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; + pB_panel += (size_t)16 * K; + pB_descales_panel += (size_t)16 * block_count; + } + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)8 * k; + const float* pB_descales = pB_descales_panel + (size_t)8 * block_start; + __m256 _out0; + __m256 _out1; + __m256 _out2; + __m256 _out3; + if (k == 0) + { + _out0 = (__m256)__lasx_xvldi(0); + _out1 = (__m256)__lasx_xvldi(0); + _out2 = (__m256)__lasx_xvldi(0); + _out3 = (__m256)__lasx_xvldi(0); + } + else + { + _out0 = (__m256)__lasx_xvld(outptr, 0); + _out1 = (__m256)__lasx_xvld(outptr + 8, 0); + _out2 = (__m256)__lasx_xvld(outptr + 16, 0); + _out3 = (__m256)__lasx_xvld(outptr + 24, 0); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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); + __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_kk0) + { + __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_kk0) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _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))); + _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, 0); + __lasx_xvst(_out1, outptr + 8, 0); + __lasx_xvst(_out2, outptr + 16, 0); + __lasx_xvst(_out3, outptr + 24, 0); + outptr += 32; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __loongarch_asx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + __m128 _out00; + __m128 _out01; + __m128 _out10; + __m128 _out11; + __m128 _out20; + __m128 _out21; + __m128 _out30; + __m128 _out31; + if (k == 0) + { + _out00 = (__m128)__lsx_vldi(0); + _out01 = (__m128)__lsx_vldi(0); + _out10 = (__m128)__lsx_vldi(0); + _out11 = (__m128)__lsx_vldi(0); + _out20 = (__m128)__lsx_vldi(0); + _out21 = (__m128)__lsx_vldi(0); + _out30 = (__m128)__lsx_vldi(0); + _out31 = (__m128)__lsx_vldi(0); + } + else + { + _out00 = (__m128)__lsx_vld(outptr, 0); + _out01 = (__m128)__lsx_vld(outptr + 4, 0); + _out10 = (__m128)__lsx_vld(outptr + 8, 0); + _out11 = (__m128)__lsx_vld(outptr + 12, 0); + _out20 = (__m128)__lsx_vld(outptr + 16, 0); + _out21 = (__m128)__lsx_vld(outptr + 20, 0); + _out30 = (__m128)__lsx_vld(outptr + 24, 0); + _out31 = (__m128)__lsx_vld(outptr + 28, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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 _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(_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, _pB0), _pA1, _pB0); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vhaddw_w_h(_s, _s)); + _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(_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_kk0) + { + __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_kk0) + { + __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 _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(_pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA1, _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA1, _pB1); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA2, _pB0); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA2, _pB1); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA3, _pB0); + _sum30 = __lsx_vadd_w(_sum30, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA3, _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; + } + __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_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + __m128 _out0; + __m128 _out1; + __m128 _out2; + __m128 _out3; + if (k == 0) + { + _out0 = (__m128)__lsx_vldi(0); + _out1 = (__m128)__lsx_vldi(0); + _out2 = (__m128)__lsx_vldi(0); + _out3 = (__m128)__lsx_vldi(0); + } + else + { + _out0 = (__m128)__lsx_vld(outptr, 0); + _out1 = (__m128)__lsx_vld(outptr + 4, 0); + _out2 = (__m128)__lsx_vld(outptr + 8, 0); + _out3 = (__m128)__lsx_vld(outptr + 12, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA = __lsx_vld(pA, 0); + __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_kk0) + { + __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_kk0) + { + __m128i _pB = __lsx_vldrepl_w(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _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)); + _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, 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_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + __m128 _out0; + __m128 _out1; + if (k == 0) + { + _out0 = (__m128)__lsx_vldi(0); + _out1 = (__m128)__lsx_vldi(0); + } + else + { + __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 kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + __m128i _sum1 = __lsx_vreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __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); + __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)); + __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)); + _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)); + _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_kk0) + { + __m128i _pA = __lsx_vldrepl_w(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _pB0 = __lsx_vldrepl_h(pB, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + __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); + __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, 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_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + __m128 _out0; + if (k == 0) + { + _out0 = (__m128)__lsx_vldi(0); + } + else + { + _out0 = (__m128)__lsx_vld(outptr, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __m128i _pA = __lsx_vldrepl_d(pA, 0); + __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; + pB += 2; + kk += 2; + } + if (kk < max_kk0) + { + __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++; + } + __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_vst((__m128i)_out0, outptr, 0); + outptr += 4; + pB_panel += K; + pB_descales_panel += block_count; + } + pAT += A_hstep * 4; + pAT_descales += A_descales_hstep * 4; + } +#endif // __loongarch_sx + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __loongarch_sx +#if __loongarch_asx + for (; jj + 15 < max_jj; jj += 16) + { + const signed char* pB0 = pB_panel + (size_t)8 * k; + const signed char* pB1 = pB_panel + (size_t)8 * K + (size_t)8 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)8 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)8 * block_count + (size_t)8 * block_start; + __m256 _out00; + __m256 _out01; + __m256 _out10; + __m256 _out11; + if (k == 0) + { + _out00 = (__m256)__lasx_xvldi(0); + _out01 = (__m256)__lasx_xvldi(0); + _out10 = (__m256)__lasx_xvldi(0); + _out11 = (__m256)__lasx_xvldi(0); + } + else + { + _out00 = (__m256)__lasx_xvld(outptr, 0); + _out01 = (__m256)__lasx_xvld(outptr + 8, 0); + _out10 = (__m256)__lasx_xvld(outptr + 16, 0); + _out11 = (__m256)__lasx_xvld(outptr + 24, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __m256i _pB0 = __lasx_xvld(pB0, 0); + __m256i _pB1 = __lasx_xvld(pB1, 0); + __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); + _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_kk0) + { + __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_kk0) + { + __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 _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(_pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA1, _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA1, _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; + } + __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; + pB_panel += (size_t)16 * K; + pB_descales_panel += (size_t)16 * block_count; + } + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)8 * k; + const float* pB_descales = pB_descales_panel + (size_t)8 * block_start; + __m256 _out0; + __m256 _out1; + if (k == 0) + { + _out0 = (__m256)__lasx_xvldi(0); + _out1 = (__m256)__lasx_xvldi(0); + } + else + { + _out0 = (__m256)__lasx_xvld(outptr, 0); + _out1 = (__m256)__lasx_xvld(outptr + 8, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m256i _sum0 = __lasx_xvreplgr2vr_w(0); + __m256i _sum1 = __lasx_xvreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __m256i _pB = __lasx_xvld(pB, 0); + __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; + } + if (kk + 1 < max_kk0) + { + __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_kk0) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _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; + 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, 0); + __lasx_xvst(_out1, outptr + 8, 0); + outptr += 16; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __loongarch_asx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + __m128 _out00; + __m128 _out01; + __m128 _out10; + __m128 _out11; + if (k == 0) + { + _out00 = (__m128)__lsx_vldi(0); + _out01 = (__m128)__lsx_vldi(0); + _out10 = (__m128)__lsx_vldi(0); + _out11 = (__m128)__lsx_vldi(0); + } + else + { + _out00 = (__m128)__lsx_vld(outptr, 0); + _out01 = (__m128)__lsx_vld(outptr + 4, 0); + _out10 = (__m128)__lsx_vld(outptr + 8, 0); + _out11 = (__m128)__lsx_vld(outptr + 12, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += 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_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pB0 = __lsx_vld(pB0, 0); + __m128i _pB1 = __lsx_vld(pB1, 0); + __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); + _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_kk0) + { + __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_kk0) + { + __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 _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(_pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA1, _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA1, _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; + } + __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; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + __m128 _out0; + __m128 _out1; + if (k == 0) + { + _out0 = (__m128)__lsx_vldi(0); + _out1 = (__m128)__lsx_vldi(0); + } + else + { + _out0 = (__m128)__lsx_vld(outptr, 0); + _out1 = (__m128)__lsx_vld(outptr + 4, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + __m128i _sum1 = __lsx_vreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pB = __lsx_vld(pB, 0); + __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; + } + if (kk + 1 < max_kk0) + { + __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_kk0) + { + __m128i _pB = __lsx_vldrepl_w(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _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; + 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, 0); + __lsx_vst((__m128i)_out1, outptr + 4, 0); + outptr += 8; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // __loongarch_sx + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float out00; + float out01; + float out10; + float out11; + if (k == 0) + { + out00 = 0.f; + out01 = 0.f; + out10 = 0.f; + out11 = 0.f; + } + else + { + out00 = outptr[0]; + out01 = outptr[1]; + out10 = outptr[2]; + out11 = outptr[3]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum00 = 0; + int sum01 = 0; + int sum10 = 0; + int sum11 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + 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_kk0) + { + 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_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float out0; + float out1; + if (k == 0) + { + out0 = 0.f; + out1 = 0.f; + } + else + { + out0 = outptr[0]; + out1 = outptr[1]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum0 = 0; + int sum1 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + 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_kk0) + { + 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_panel += K; + pB_descales_panel += block_count; + } + pAT += A_hstep * 2; + pAT_descales += A_descales_hstep * 2; + } + for (; ii < max_ii; ii++) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __loongarch_sx +#if __loongarch_asx + for (; jj + 15 < max_jj; jj += 16) + { + const signed char* pB0 = pB_panel + (size_t)8 * k; + const signed char* pB1 = pB_panel + (size_t)8 * K + (size_t)8 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)8 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)8 * block_count + (size_t)8 * block_start; + __m256 _out00; + __m256 _out01; + if (k == 0) + { + _out00 = (__m256)__lasx_xvldi(0); + _out01 = (__m256)__lasx_xvldi(0); + } + else + { + _out00 = (__m256)__lasx_xvld(outptr, 0); + _out01 = (__m256)__lasx_xvld(outptr + 8, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m256i _sum00 = __lasx_xvreplgr2vr_w(0); + __m256i _sum01 = __lasx_xvreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __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_kk0) + { + __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; + } + __lasx_xvst(_out00, outptr, 0); + __lasx_xvst(_out01, outptr + 8, 0); + outptr += 16; + pB_panel += (size_t)16 * K; + pB_descales_panel += (size_t)16 * block_count; + } + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)8 * k; + const float* pB_descales = pB_descales_panel + (size_t)8 * block_start; + __m256 _out0; + if (k == 0) + { + _out0 = (__m256)__lasx_xvldi(0); + } + else + { + _out0 = (__m256)__lasx_xvld(outptr, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m256i _sum0 = __lasx_xvreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __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_kk0) + { + __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, 0); + outptr += 8; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __loongarch_asx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + __m128 _out00; + __m128 _out01; + if (k == 0) + { + _out00 = (__m128)__lsx_vldi(0); + _out01 = (__m128)__lsx_vldi(0); + } + else + { + _out00 = (__m128)__lsx_vld(outptr, 0); + _out01 = (__m128)__lsx_vld(outptr + 4, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m128i _sum00 = __lsx_vreplgr2vr_w(0); + __m128i _sum01 = __lsx_vreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __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_kk0) + { + __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; + } + __lsx_vst((__m128i)_out00, outptr, 0); + __lsx_vst((__m128i)_out01, outptr + 4, 0); + outptr += 8; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + __m128 _out0; + if (k == 0) + { + _out0 = (__m128)__lsx_vldi(0); + } + else + { + _out0 = (__m128)__lsx_vld(outptr, 0); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + __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_kk0) + { + __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, 0); + outptr += 4; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // __loongarch_sx + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float out0; + float out1; + if (k == 0) + { + out0 = 0.f; + out1 = 0.f; + } + else + { + out0 = outptr[0]; + out1 = outptr[1]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum0 = 0; + int sum1 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + 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_kk0) + { + 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[0] = out0; + outptr[1] = out1; + outptr += 2; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float out0; + if (k == 0) + { + out0 = 0.f; + } + else + { + out0 = outptr[0]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + int sum0 = 0; + const int max_kk0 = std::min(max_kk - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + asm volatile("" + : + : + : "memory"); + + 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_kk0) + { + sum0 += pA[0] * pB[0] + pA[1] * pB[1]; + pB += 2; + pA += 2; + kk += 2; + } + if (kk < max_kk0) + { + sum0 += pA[0] * pB[0]; + pB++; + pA++; + } + out0 += sum0 * *pA_descales++ * *pB_descales++; + } + *outptr++ = out0; + pB_panel += K; + pB_descales_panel += block_count; + } + 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, float alpha, float beta) +{ + const float* pp = topT; + 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; + const float* pC_base = C; + float* outptr = (float*)top_blob + (size_t)i * out_hstep + j; + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + float* p0 = outptr; + float* p1 = p0 + out_hstep; + float* p2 = p1 + out_hstep; + float* p3 = p2 + out_hstep; + float* p4 = p3 + out_hstep; + float* p5 = p4 + out_hstep; + float* p6 = p5 + out_hstep; + float* p7 = p6 + out_hstep; + + 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 + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + for (; jj + 7 < max_jj; jj += 8) + { + __m256i _r0 = __lasx_xvld(pp, 0); + __m256i _r1 = __lasx_xvld(pp + 8, 0); + __m256i _r2 = __lasx_xvld(pp + 16, 0); + __m256i _r3 = __lasx_xvld(pp + 24, 0); + __m256i _r4 = __lasx_xvld(pp + 32, 0); + __m256i _r5 = __lasx_xvld(pp + 40, 0); + __m256i _r6 = __lasx_xvld(pp + 48, 0); + __m256i _r7 = __lasx_xvld(pp + 56, 0); + pp += 64; + __m256i _tmp0 = _r0; + __m256i _tmp1 = __lasx_xvshuf4i_w(_r1, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp2 = _r2; + __m256i _tmp3 = __lasx_xvshuf4i_w(_r3, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp4 = _r4; + __m256i _tmp5 = __lasx_xvshuf4i_w(_r5, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp6 = _r6; + __m256i _tmp7 = __lasx_xvshuf4i_w(_r7, _LSX_SHUFFLE(2, 1, 0, 3)); + _r0 = __lasx_xvilvl_w(_tmp3, _tmp0); + _r1 = __lasx_xvilvh_w(_tmp3, _tmp0); + _r2 = __lasx_xvilvl_w(_tmp1, _tmp2); + _r3 = __lasx_xvilvh_w(_tmp1, _tmp2); + _r4 = __lasx_xvilvl_w(_tmp7, _tmp4); + _r5 = __lasx_xvilvh_w(_tmp7, _tmp4); + _r6 = __lasx_xvilvl_w(_tmp5, _tmp6); + _r7 = __lasx_xvilvh_w(_tmp5, _tmp6); + _tmp0 = __lasx_xvilvl_d(_r2, _r0); + _tmp1 = __lasx_xvilvh_d(_r2, _r0); + _tmp2 = __lasx_xvilvl_d(_r1, _r3); + _tmp3 = __lasx_xvilvh_d(_r1, _r3); + _tmp4 = __lasx_xvilvl_d(_r6, _r4); + _tmp5 = __lasx_xvilvh_d(_r6, _r4); + _tmp6 = __lasx_xvilvl_d(_r5, _r7); + _tmp7 = __lasx_xvilvh_d(_r5, _r7); + _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)); + if (pC) + { + if (broadcast_type_C == 0) + { + __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) + { + __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); + pC += 8; + 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); + pC += 8; + 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; + } +#endif // __loongarch_asx + __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + __m128 _beta128 = __lsx_vreplfr2vr_s(beta); + for (; jj + 3 < max_jj; jj += 4) + { + __m128i _r0 = __lsx_vld(pp, 0); + __m128i _r1 = __lsx_vld(pp + 8, 0); + __m128i _r2 = __lsx_vld(pp + 16, 0); + __m128i _r3 = __lsx_vld(pp + 24, 0); + _r2 = __lsx_vshuf4i_w(_r2, _LSX_SHUFFLE(1, 0, 3, 2)); + _r3 = __lsx_vshuf4i_w(_r3, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_r0, _r1, _r2, _r3); + _r1 = __lsx_vshuf4i_w(_r1, _LSX_SHUFFLE(2, 1, 0, 3)); + _r2 = __lsx_vshuf4i_w(_r2, _LSX_SHUFFLE(1, 0, 3, 2)); + _r3 = __lsx_vshuf4i_w(_r3, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _r4 = __lsx_vld(pp + 4, 0); + __m128i _r5 = __lsx_vld(pp + 12, 0); + __m128i _r6 = __lsx_vld(pp + 20, 0); + __m128i _r7 = __lsx_vld(pp + 28, 0); + pp += 32; + _r6 = __lsx_vshuf4i_w(_r6, _LSX_SHUFFLE(1, 0, 3, 2)); + _r7 = __lsx_vshuf4i_w(_r7, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_r4, _r5, _r6, _r7); + _r5 = __lsx_vshuf4i_w(_r5, _LSX_SHUFFLE(2, 1, 0, 3)); + _r6 = __lsx_vshuf4i_w(_r6, _LSX_SHUFFLE(1, 0, 3, 2)); + _r7 = __lsx_vshuf4i_w(_r7, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128 _f0 = (__m128)_r0; + __m128 _f1 = (__m128)_r1; + __m128 _f2 = (__m128)_r2; + __m128 _f3 = (__m128)_r3; + __m128 _f4 = (__m128)_r4; + __m128 _f5 = (__m128)_r5; + __m128 _f6 = (__m128)_r6; + __m128 _f7 = (__m128)_r7; + 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); + _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) + { + __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); + _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); + } + pC += 4; + } + 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); + pC += 4; + } + } + 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; + } + 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); + pp += 16; + __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; + 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); + _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) + { + __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); + _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); + } + pC += 2; + } + 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); + pC += 2; + } + } + 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; + } + 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) + { + __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); + } + pC++; + } + if (broadcast_type_C == 4) + { + __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) + { + _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++; + } + outptr += out_hstep * 8; + } + for (; ii + 3 < max_ii; ii += 4) + { + float* p0 = outptr; + float* p1 = p0 + out_hstep; + float* p2 = p1 + out_hstep; + float* p3 = p2 + out_hstep; + + 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 + __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, 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) + { + __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) + { + __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); + _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) + { + __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); + _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); + } + pC += 16; + } + 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); + pC += 16; + } + } + 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; + } + for (; jj + 7 < max_jj; jj += 8) + { + __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) + { + _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); + } + pC += 8; + } + 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); + pC += 8; + } + } + 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; + } +#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, 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) + { + __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) + { + __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); + _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) + { + __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); + _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); + } + pC += 8; + } + 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); + pC += 8; + } + } + 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; + } + for (; jj + 3 < max_jj; jj += 4) + { + __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) + { + _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); + } + pC += 4; + } + 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); + pC += 4; + } + } + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + __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) + { + __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); + } + pC += 2; + } + 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); + pC += 2; + } + } + 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; + } + for (; jj < max_jj; jj++) + { + __m128i _fi = __lsx_vld(pp, 0); + pp += 4; + __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); + 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); + __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++; + } + outptr += out_hstep * 4; + } +#endif // __loongarch_sx + for (; ii + 1 < max_ii; ii += 2) + { + float* p0 = outptr; + float* p1 = p0 + out_hstep; + + 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_sx +#if __loongarch_asx + __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, 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); + pp += 32; + if (pC) + { + if (broadcast_type_C == 0) + { + __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) + { + __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); + _f11 = __lasx_xvfadd_s(_f11, _c1); + } + if (broadcast_type_C == 3) + { + __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); + _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); + } + pC += 16; + } + 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); + pC += 16; + } + } + 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; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); + pp += 16; + 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); + } + pC += 8; + } + 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); + pC += 8; + } + } + 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; + } +#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, 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); + pp += 16; + if (pC) + { + if (broadcast_type_C == 0) + { + __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) + { + __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); + _f11 = __lsx_vfadd_s(_f11, _c1); + } + 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); + 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); + } + pC += 8; + } + 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); + pC += 8; + } + } + 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; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); + pp += 8; + 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); + } + pC += 4; + } + 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); + pC += 4; + } + } + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp, 0); + __m128 _f1 = (__m128)__lsx_vldrepl_d(pp + 2, 0); + pp += 4; + 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); + } + pC += 2; + } + 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); + pC += 2; + } + } + 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; + } +#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]; + pp += 4; + 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; + } + for (; jj < max_jj; jj++) + { + float f0 = pp[0]; + float f1 = pp[1]; + pp += 2; + 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; + } + 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) + { + f0 *= alpha; + f1 *= alpha; + } + p0[0] = f0; + p1[0] = f1; + p0++; + p1++; + } + outptr += out_hstep * 2; + } + for (; ii < max_ii; ii++) + { + float* p0 = outptr; + + 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_sx +#if __loongarch_asx + __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, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); + pp += 16; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + __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) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + pC += 16; + 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; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + pp += 8; + 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); + pC += 8; + 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; + } +#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, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); + pp += 8; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + __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) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + pC += 8; + 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; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + pp += 4; + 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); + pC += 4; + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp, 0); + pp += 2; + 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); + pC += 2; + 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; + } +#endif // __loongarch_sx + for (; jj < 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++; + } + outptr += out_hstep; + } +} + +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, float alpha, float beta) +{ + const float* pp = topT; + 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; + const float* pC_base = C; + float* outptr = (float*)top_blob + (size_t)j * out_hstep + i; + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + float* p0 = outptr; + 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) + { + __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; + + __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + __m128 _beta = __lsx_vreplfr2vr_s(beta); + int jj = 0; +#if __loongarch_asx + __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 _r0 = __lasx_xvld(pp, 0); + __m256i _r1 = __lasx_xvld(pp + 8, 0); + __m256i _r2 = __lasx_xvld(pp + 16, 0); + __m256i _r3 = __lasx_xvld(pp + 24, 0); + __m256i _r4 = __lasx_xvld(pp + 32, 0); + __m256i _r5 = __lasx_xvld(pp + 40, 0); + __m256i _r6 = __lasx_xvld(pp + 48, 0); + __m256i _r7 = __lasx_xvld(pp + 56, 0); + pp += 64; + __m256i _tmp0 = _r0; + __m256i _tmp1 = __lasx_xvshuf4i_w(_r1, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp2 = _r2; + __m256i _tmp3 = __lasx_xvshuf4i_w(_r3, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp4 = _r4; + __m256i _tmp5 = __lasx_xvshuf4i_w(_r5, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp6 = _r6; + __m256i _tmp7 = __lasx_xvshuf4i_w(_r7, _LSX_SHUFFLE(2, 1, 0, 3)); + _r0 = __lasx_xvilvl_w(_tmp3, _tmp0); + _r1 = __lasx_xvilvh_w(_tmp3, _tmp0); + _r2 = __lasx_xvilvl_w(_tmp1, _tmp2); + _r3 = __lasx_xvilvh_w(_tmp1, _tmp2); + _r4 = __lasx_xvilvl_w(_tmp7, _tmp4); + _r5 = __lasx_xvilvh_w(_tmp7, _tmp4); + _r6 = __lasx_xvilvl_w(_tmp5, _tmp6); + _r7 = __lasx_xvilvh_w(_tmp5, _tmp6); + _tmp0 = __lasx_xvilvl_d(_r2, _r0); + _tmp1 = __lasx_xvilvh_d(_r2, _r0); + _tmp2 = __lasx_xvilvl_d(_r1, _r3); + _tmp3 = __lasx_xvilvh_d(_r1, _r3); + _tmp4 = __lasx_xvilvl_d(_r6, _r4); + _tmp5 = __lasx_xvilvh_d(_r6, _r4); + _tmp6 = __lasx_xvilvl_d(_r5, _r7); + _tmp7 = __lasx_xvilvh_d(_r5, _r7); + _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)); + 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); + } + pC += 8; + } + 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)); + pC += 8; + } + } + 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 + out_hstep, 0); + __lasx_xvst(_f2, p0 + out_hstep * 2, 0); + __lasx_xvst(_f3, p0 + out_hstep * 3, 0); + __lasx_xvst(_f4, p0 + out_hstep * 4, 0); + __lasx_xvst(_f5, p0 + out_hstep * 5, 0); + __lasx_xvst(_f6, p0 + out_hstep * 6, 0); + __lasx_xvst(_f7, p0 + out_hstep * 7, 0); + p0 += out_hstep * 8; + } +#endif // __loongarch_asx + for (; jj + 3 < max_jj; jj += 4) + { + __m128i _r0 = __lsx_vld(pp, 0); + __m128i _r1 = __lsx_vld(pp + 8, 0); + __m128i _r2 = __lsx_vld(pp + 16, 0); + __m128i _r3 = __lsx_vld(pp + 24, 0); + _r2 = __lsx_vshuf4i_w(_r2, _LSX_SHUFFLE(1, 0, 3, 2)); + _r3 = __lsx_vshuf4i_w(_r3, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_r0, _r1, _r2, _r3); + _r1 = __lsx_vshuf4i_w(_r1, _LSX_SHUFFLE(2, 1, 0, 3)); + _r2 = __lsx_vshuf4i_w(_r2, _LSX_SHUFFLE(1, 0, 3, 2)); + _r3 = __lsx_vshuf4i_w(_r3, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _r4 = __lsx_vld(pp + 4, 0); + __m128i _r5 = __lsx_vld(pp + 12, 0); + __m128i _r6 = __lsx_vld(pp + 20, 0); + __m128i _r7 = __lsx_vld(pp + 28, 0); + pp += 32; + _r6 = __lsx_vshuf4i_w(_r6, _LSX_SHUFFLE(1, 0, 3, 2)); + _r7 = __lsx_vshuf4i_w(_r7, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_r4, _r5, _r6, _r7); + _r5 = __lsx_vshuf4i_w(_r5, _LSX_SHUFFLE(2, 1, 0, 3)); + _r6 = __lsx_vshuf4i_w(_r6, _LSX_SHUFFLE(1, 0, 3, 2)); + _r7 = __lsx_vshuf4i_w(_r7, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128 _f0 = (__m128)_r0; + __m128 _f1 = (__m128)_r1; + __m128 _f2 = (__m128)_r2; + __m128 _f3 = (__m128)_r3; + __m128 _f4 = (__m128)_r4; + __m128 _f5 = (__m128)_r5; + __m128 _f6 = (__m128)_r6; + __m128 _f7 = (__m128)_r7; + 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); + } + pC += 4; + } + 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)); + pC += 4; + } + } + 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 + out_hstep, 0); + __lsx_vst((__m128i)_f5, p0 + out_hstep + 4, 0); + __lsx_vst((__m128i)_f2, p0 + out_hstep * 2, 0); + __lsx_vst((__m128i)_f6, p0 + out_hstep * 2 + 4, 0); + __lsx_vst((__m128i)_f3, p0 + out_hstep * 3, 0); + __lsx_vst((__m128i)_f7, p0 + out_hstep * 3 + 4, 0); + p0 += out_hstep * 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); + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __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) + { + _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 + out_hstep, 0); + __lsx_vst((__m128i)_f3, p0 + out_hstep + 4, 0); + p0 += out_hstep * 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); + } + pC++; + } + if (broadcast_type_C == 4) + { + __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) + { + _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 += out_hstep; + } + outptr += 8; + } + for (; ii + 3 < max_ii; ii += 4) + { + float* p0 = outptr; + 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; + + __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + __m128 _beta = __lsx_vreplfr2vr_s(beta); + int jj = 0; +#if __loongarch_asx + __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, 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) + { + __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) + { + __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); + _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) + { + __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); + _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); + } + pC += 16; + } + 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); + pC += 16; + } + } + 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 + out_hstep, 0); + __lsx_vst(__lasx_extract_128_lo(_r2), p0 + out_hstep * 2, 0); + __lsx_vst(__lasx_extract_128_lo(_r3), p0 + out_hstep * 3, 0); + __lsx_vst(__lasx_extract_128_hi(_r0), p0 + out_hstep * 4, 0); + __lsx_vst(__lasx_extract_128_hi(_r1), p0 + out_hstep * 5, 0); + __lsx_vst(__lasx_extract_128_hi(_r2), p0 + out_hstep * 6, 0); + __lsx_vst(__lasx_extract_128_hi(_r3), p0 + out_hstep * 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 + out_hstep * 8, 0); + __lsx_vst(__lasx_extract_128_lo(_r1), p0 + out_hstep * 9, 0); + __lsx_vst(__lasx_extract_128_lo(_r2), p0 + out_hstep * 10, 0); + __lsx_vst(__lasx_extract_128_lo(_r3), p0 + out_hstep * 11, 0); + __lsx_vst(__lasx_extract_128_hi(_r0), p0 + out_hstep * 12, 0); + __lsx_vst(__lasx_extract_128_hi(_r1), p0 + out_hstep * 13, 0); + __lsx_vst(__lasx_extract_128_hi(_r2), p0 + out_hstep * 14, 0); + __lsx_vst(__lasx_extract_128_hi(_r3), p0 + out_hstep * 15, 0); + p0 += out_hstep * 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __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) + { + __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); + } + pC += 8; + } + 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); + pC += 8; + } + } + 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 + out_hstep, 0); + __lsx_vst(__lasx_extract_128_lo(_r2), p0 + out_hstep * 2, 0); + __lsx_vst(__lasx_extract_128_lo(_r3), p0 + out_hstep * 3, 0); + __lsx_vst(__lasx_extract_128_hi(_r0), p0 + out_hstep * 4, 0); + __lsx_vst(__lasx_extract_128_hi(_r1), p0 + out_hstep * 5, 0); + __lsx_vst(__lasx_extract_128_hi(_r2), p0 + out_hstep * 6, 0); + __lsx_vst(__lasx_extract_128_hi(_r3), p0 + out_hstep * 7, 0); + p0 += out_hstep * 8; + } +#endif // __loongarch_asx + for (; jj + 7 < max_jj; jj += 8) + { + __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) + { + 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); + } + pC += 8; + } + 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)); + pC += 8; + } + } + 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 + out_hstep, 0); + __lsx_vst((__m128i)_f20, p0 + out_hstep * 2, 0); + __lsx_vst((__m128i)_f30, p0 + out_hstep * 3, 0); + __lsx_vst((__m128i)_f01, p0 + out_hstep * 4, 0); + __lsx_vst((__m128i)_f11, p0 + out_hstep * 5, 0); + __lsx_vst((__m128i)_f21, p0 + out_hstep * 6, 0); + __lsx_vst((__m128i)_f31, p0 + out_hstep * 7, 0); + p0 += out_hstep * 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __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) + { + 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); + } + pC += 4; + } + 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)); + pC += 4; + } + } + 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 + out_hstep, 0); + __lsx_vst((__m128i)_f2, p0 + out_hstep * 2, 0); + __lsx_vst((__m128i)_f3, p0 + out_hstep * 3, 0); + p0 += out_hstep * 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __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); + __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); + __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); + _f1 = __lsx_vfadd_s(_f1, _cc1); + } + else + { + _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) + { + _f0 = __lsx_vfmul_s(_f0, _alpha); + _f1 = __lsx_vfmul_s(_f1, _alpha); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p0 + out_hstep, 0); + p0 += out_hstep * 2; + } + for (; jj < max_jj; jj++) + { + __m128i _fi = __lsx_vld(pp, 0); + pp += 4; + __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); + pC++; + } + 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); + pC++; + } + } + if (alpha != 1.f) + _f = __lsx_vfmul_s(_f, _alpha); + __lsx_vst((__m128i)_f, p0, 0); + p0 += out_hstep; + } + outptr += 4; + } +#endif // __loongarch_sx + for (; ii + 1 < max_ii; ii += 2) + { + float* p0 = outptr; + 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)); + } + __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + __m128 _beta = __lsx_vreplfr2vr_s(beta); +#if __loongarch_asx + __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, 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); + pp += 32; + if (pC) + { + if (broadcast_type_C == 0) + { + __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) + { + __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); + _f11 = __lasx_xvfadd_s(_f11, _c1); + } + if (broadcast_type_C == 3) + { + __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); + _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); + } + pC += 16; + } + 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); + pC += 16; + } + } + 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 + out_hstep, 0, 1); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 2, 0, 0); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 3, 0, 1); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 4, 0, 2); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 5, 0, 3); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 6, 0, 2); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 7, 0, 3); + _tmp0 = __lasx_xvilvl_w((__m256i)_f11, (__m256i)_f01); + _tmp1 = __lasx_xvilvh_w((__m256i)_f11, (__m256i)_f01); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 8, 0, 0); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 9, 0, 1); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 10, 0, 0); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 11, 0, 1); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 12, 0, 2); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 13, 0, 3); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 14, 0, 2); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 15, 0, 3); + p0 += out_hstep * 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); + pp += 16; + 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); + } + pC += 8; + } + 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); + pC += 8; + } + } + 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 + out_hstep, 0, 1); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 2, 0, 0); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 3, 0, 1); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 4, 0, 2); + __lasx_xvstelm_d(_tmp0, p0 + out_hstep * 5, 0, 3); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 6, 0, 2); + __lasx_xvstelm_d(_tmp1, p0 + out_hstep * 7, 0, 3); + p0 += out_hstep * 8; + } +#endif // __loongarch_asx + for (; jj + 7 < max_jj; jj += 8) + { + __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); + pp += 16; + if (pC) + { + if (broadcast_type_C == 0) + { + __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) + { + __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); + _f11 = __lsx_vfadd_s(_f11, _c1); + } + 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); + 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); + } + pC += 8; + } + 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); + pC += 8; + } + } + 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 + out_hstep, 0, 1); + __lsx_vstelm_d(_tmp1, p0 + out_hstep * 2, 0, 0); + __lsx_vstelm_d(_tmp1, p0 + out_hstep * 3, 0, 1); + _tmp0 = __lsx_vilvl_w((__m128i)_f11, (__m128i)_f01); + _tmp1 = __lsx_vilvh_w((__m128i)_f11, (__m128i)_f01); + __lsx_vstelm_d(_tmp0, p0 + out_hstep * 4, 0, 0); + __lsx_vstelm_d(_tmp0, p0 + out_hstep * 5, 0, 1); + __lsx_vstelm_d(_tmp1, p0 + out_hstep * 6, 0, 0); + __lsx_vstelm_d(_tmp1, p0 + out_hstep * 7, 0, 1); + p0 += out_hstep * 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); + pp += 8; + 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); + } + pC += 4; + } + 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); + pC += 4; + } + } + 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 + out_hstep, 0, 1); + __lsx_vstelm_d(_tmp1, p0 + out_hstep * 2, 0, 0); + __lsx_vstelm_d(_tmp1, p0 + out_hstep * 3, 0, 1); + p0 += out_hstep * 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f = (__m128)__lsx_vshuf4i_w(__lsx_vld(pp, 0), _LSX_SHUFFLE(3, 1, 2, 0)); + pp += 4; + __m128i _r0; + __m128i _r1; + 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); + __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) + { + __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); + pC += 2; + } + } + if (alpha != 1.f) + _f = __lsx_vfmul_s(_f, _alpha); + __lsx_vstelm_d((__m128i)_f, p0, 0, 0); + __lsx_vstelm_d((__m128i)_f, p0 + out_hstep, 0, 1); + p0 += out_hstep * 2; + } + for (; jj < max_jj; jj++) + { + __m128i _fi = __lsx_vldrepl_d(pp, 0); + pp += 2; + __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); + pC++; + } + 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); + pC++; + } + } + if (alpha != 1.f) + _f = __lsx_vfmul_s(_f, _alpha); + __lsx_vstelm_d((__m128i)_f, p0, 0, 0); + p0 += out_hstep; + } +#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]; + pp += 4; + 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[out_hstep] = f01; + p0[out_hstep + 1] = f11; + p0 += out_hstep * 2; + } + for (; jj < max_jj; jj++) + { + float f0 = pp[0]; + float f1 = pp[1]; + pp += 2; + 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; + } + 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) + { + f0 *= alpha; + f1 *= alpha; + } + p0[0] = f0; + p0[1] = f1; + p0 += out_hstep; + } + outptr += 2; + } + for (; ii < max_ii; ii++) + { + float* p0 = outptr; + 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_sx +#if __loongarch_asx + __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(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); + pp += 16; + 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) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + pC += 16; + 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 (out_hstep == 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 + out_hstep, 0, 1); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 2, 0, 2); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 3, 0, 3); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 4, 0, 4); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 5, 0, 5); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 6, 0, 6); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 7, 0, 7); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 8, 0, 0); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 9, 0, 1); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 10, 0, 2); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 11, 0, 3); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 12, 0, 4); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 13, 0, 5); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 14, 0, 6); + __lasx_xvstelm_w((__m256i)_f1, p0 + out_hstep * 15, 0, 7); + } + p0 += out_hstep * 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + pp += 8; + 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); + pC += 8; + 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 (out_hstep == 1) + __lasx_xvst(_f0, p0, 0); + else + { + __lasx_xvstelm_w((__m256i)_f0, p0, 0, 0); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep, 0, 1); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 2, 0, 2); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 3, 0, 3); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 4, 0, 4); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 5, 0, 5); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 6, 0, 6); + __lasx_xvstelm_w((__m256i)_f0, p0 + out_hstep * 7, 0, 7); + } + p0 += out_hstep * 8; + } +#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(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); + pp += 8; + 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) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + pC += 8; + 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 (out_hstep == 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 + out_hstep, 0, 1); + __lsx_vstelm_w((__m128i)_f0, p0 + out_hstep * 2, 0, 2); + __lsx_vstelm_w((__m128i)_f0, p0 + out_hstep * 3, 0, 3); + __lsx_vstelm_w((__m128i)_f1, p0 + out_hstep * 4, 0, 0); + __lsx_vstelm_w((__m128i)_f1, p0 + out_hstep * 5, 0, 1); + __lsx_vstelm_w((__m128i)_f1, p0 + out_hstep * 6, 0, 2); + __lsx_vstelm_w((__m128i)_f1, p0 + out_hstep * 7, 0, 3); + } + p0 += out_hstep * 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + pp += 4; + 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); + pC += 4; + 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 (out_hstep == 1) + __lsx_vst((__m128i)_f0, p0, 0); + else + { + __lsx_vstelm_w((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_w((__m128i)_f0, p0 + out_hstep, 0, 1); + __lsx_vstelm_w((__m128i)_f0, p0 + out_hstep * 2, 0, 2); + __lsx_vstelm_w((__m128i)_f0, p0 + out_hstep * 3, 0, 3); + } + p0 += out_hstep * 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp, 0); + pp += 2; + 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 _cc = (__m128)__lsx_vldrepl_d(pC, 0); + pC += 2; + 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 (out_hstep == 1) + __lsx_vstelm_d((__m128i)_f0, p0, 0, 0); + else + { + __lsx_vstelm_w((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_w((__m128i)_f0, p0 + out_hstep, 0, 1); + } + p0 += out_hstep * 2; + } +#endif // __loongarch_sx + for (; jj < 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 += out_hstep; + } + outptr++; + } +} + +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(); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + + int tile_size = (int)sqrtf((float)l2_cache_size / (2 * sizeof(signed char) + sizeof(float))); + +#if __loongarch_sx + TILE_M = std::max(8, tile_size / 8 * 8); +#if __loongarch_asx + TILE_N = std::max(16, tile_size / 16 * 16); +#else + TILE_N = std::max(8, tile_size / 8 * 8); +#endif +#else + TILE_M = std::max(2, tile_size / 2 * 2); + TILE_N = std::max(2, tile_size / 2 * 2); +#endif + + TILE_K = std::max(block_size, tile_size / block_size * block_size); + + if (K > 0) + { + int nn_K = (K + TILE_K - 1) / TILE_K; + TILE_K = std::min(TILE_K, ((K + nn_K - 1) / nn_K + block_size - 1) / block_size * block_size); + TILE_K = std::min(TILE_K, K); + + if (nn_K == 1) + { + tile_size = (int)((float)l2_cache_size / 2 / sizeof(signed char) / TILE_K); + +#if __loongarch_sx + TILE_M = std::max(8, tile_size / 8 * 8); +#if __loongarch_asx + TILE_N = std::max(16, tile_size / 16 * 16); +#else + TILE_N = std::max(8, tile_size / 8 * 8); +#endif +#else + TILE_M = std::max(2, tile_size / 2 * 2); + TILE_N = std::max(2, tile_size / 2 * 2); +#endif + } + } + + TILE_M *= std::min(nT, get_physical_cpu_count()); + + if (M > 0) + { + int nn_M = (M + TILE_M - 1) / TILE_M; +#if __loongarch_sx + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + 7) / 8 * 8); +#else + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + 1) / 2 * 2); +#endif + } + + if (N > 0) + { + int nn_N = (N + TILE_N - 1) / TILE_N; +#if __loongarch_sx +#if __loongarch_asx + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 15) / 16 * 16); +#else + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 7) / 8 * 8); +#endif +#else + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 1) / 2 * 2); +#endif + } + + if (nT > 1) + { +#if __loongarch_sx + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + 7) / 8 * 8); +#else + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + 1) / 2 * 2); +#endif + } + + // always take constant TILE_M/N/K value when provided + if (constant_TILE_M > 0) + { +#if __loongarch_sx + TILE_M = (constant_TILE_M + 7) / 8 * 8; +#else + TILE_M = (constant_TILE_M + 1) / 2 * 2; +#endif + } + + if (constant_TILE_N > 0) + { +#if __loongarch_sx +#if __loongarch_asx + TILE_N = (constant_TILE_N + 15) / 16 * 16; +#else + TILE_N = (constant_TILE_N + 7) / 8 * 8; +#endif +#else + TILE_N = (constant_TILE_N + 1) / 2 * 2; +#endif + } + + 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); + } +} diff --git a/src/layer/loongarch/multiheadattention_loongarch.cpp b/src/layer/loongarch/multiheadattention_loongarch.cpp index 7e749533d1c..bc16e796b90 100644 --- a/src/layer/loongarch/multiheadattention_loongarch.cpp +++ b/src/layer/loongarch/multiheadattention_loongarch.cpp @@ -28,10 +28,362 @@ 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; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -260,13 +612,33 @@ int MultiHeadAttention_loongarch::create_pipeline(const Option& _opt) int MultiHeadAttention_loongarch::destroy_pipeline(const Option& _opt) { if (weight_block_quantize) - return 0; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 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 +649,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; } @@ -322,7 +694,17 @@ 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) - return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + } int q_blob_i = 0; int k_blob_i = 0; @@ -340,10 +722,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 +780,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 +790,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 +818,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 +869,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 +897,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 +944,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 fe13c097511..f4b729edd29 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 81f0cb108b5..51d16a3c31a 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,284 @@ 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, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + 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); + + 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; + + const int nn_MK = nn_M * nn_K; + #pragma omp parallel for num_threads(nT) + 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_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, 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); + } + + 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 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, 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); + 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 topT_tile = topT.channel(get_omp_thread_num()); + 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); + + 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); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + } + } + } + + return 0; +} + +int Gemm_mips::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; + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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); + if (ret != 0) + return ret; + if (BT_data_packed.empty() || BT_data_packed_descales.empty()) + return -100; + + 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_mips::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 || (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; + } + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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) + { + if (output_N1M) + top_blob.create(M, 1, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + } + else + { + if (output_N1M) + top_blob.create(N, 1, M, (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, 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_mips::create_pipeline(const Option& opt) { AT_data.release(); @@ -4479,6 +4761,18 @@ int Gemm_mips::create_pipeline(const Option& opt) if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + int weight_bits; + int block_size; + bool has_input_scale; + int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits == 8) + return create_pipeline_wq_int8(opt); +#endif + return 0; } @@ -4616,10 +4910,32 @@ int Gemm_mips::create_pipeline(const Option& opt) return 0; } +int Gemm_mips::destroy_pipeline(const Option& /*opt*/) +{ +#if NCNN_WEIGHT_QUANT + BT_data_wq_int8.release(); + BT_data_wq_int8_descales.release(); +#endif + + return 0; +} + 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 + int weight_bits; + int block_size; + bool has_input_scale; + int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits == 8) + return forward_wq_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 ff043262c13..b1d5c03c333 100644 --- a/src/layer/mips/gemm_mips.h +++ b/src/layer/mips/gemm_mips.h @@ -15,9 +15,15 @@ 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 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_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 +39,10 @@ class Gemm_mips : 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/mips/gemm_mips_mmi.cpp b/src/layer/mips/gemm_mips_mmi.cpp index 09bae1c3979..06696bf0cfb 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 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, 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 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, 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 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, max_kk, K, 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 00000000000..e62047cc5d7 --- /dev/null +++ b/src/layer/mips/gemm_wq_int8.h @@ -0,0 +1,6018 @@ +// 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 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 +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; + const int j8 = j; + j += nn8 * 8; + const int nn4 = (N - j) / 4; + const int j4 = j; + j += nn4 * 4; +#endif + const int nn2 = (N - j) / 2; + const int j2 = j; + j += nn2 * 2; + const int nn1 = N - j; + const int j1 = j; + + #pragma omp parallel num_threads(opt.num_threads) + { +#if __mips_msa + #pragma omp for + for (int ppj = 0; ppj < nn8; ppj++) + { + const int j = j8 + 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) + { + 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) + { + 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) + { + 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[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; + } + } + } + #pragma omp for + for (int ppj = 0; ppj < nn4; ppj++) + { + const int j = j4 + 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) + { + 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) + { + 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[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; + } + } +#endif // __mips_msa + + #pragma omp for + 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; + 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) + { + 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) + { + 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[0]; + pp[1] = p1[0]; + pp += 2; + p0++; + p1++; + } + + *pd++ = 1.f / *s0++; + *pd++ = 1.f / *s1++; + } + } + #pragma omp for + for (int ppj = 0; ppj < nn1; ppj++) + { + const int j = j1 + ppj; + 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); + 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; + 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++; + } + + *pd++ = 1.f / *s0++; + } + } + } + + 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 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, k, max_kk, block_size, input_scale_ptr); + return; + } +#endif + + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + 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 = 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_kk0 = std::min(max_kk - 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); + 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); + + 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_kk0; kk += 4) + { + 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); + _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)); + 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); + 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_kk0; kk++) + { + 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; + 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; + + 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; + 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); + v4f32 _scale4 = __msa_fill_w_f32(scale4); + v4f32 _scale5 = __msa_fill_w_f32(scale5); + v4f32 _scale6 = __msa_fill_w_f32(scale6); + v4f32 _scale7 = __msa_fill_w_f32(scale7); + kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + v4f32 _p0 = (v4f32)__msa_ld_w(p0q, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1q, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2q, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3q, 0); + v4f32 _p4 = (v4f32)__msa_ld_w(p4q, 0); + v4f32 _p5 = (v4f32)__msa_ld_w(p5q, 0); + v4f32 _p6 = (v4f32)__msa_ld_w(p6q, 0); + v4f32 _p7 = (v4f32)__msa_ld_w(p7q, 0); + if (psq) + { + v4f32 _s = (v4f32)__msa_ld_w(psq, 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); + } + _p0 = __msa_fmul_w(_p0, _scale0); + _p1 = __msa_fmul_w(_p1, _scale1); + _p2 = __msa_fmul_w(_p2, _scale2); + _p3 = __msa_fmul_w(_p3, _scale3); + _p4 = __msa_fmul_w(_p4, _scale4); + _p5 = __msa_fmul_w(_p5, _scale5); + _p6 = __msa_fmul_w(_p6, _scale6); + _p7 = __msa_fmul_w(_p7, _scale7); + + ((int64_t*)pp)[0] = float2int8(_p0, _p1); + ((int64_t*)pp)[1] = float2int8(_p2, _p3); + ((int64_t*)pp)[2] = float2int8(_p4, _p5); + ((int64_t*)pp)[3] = float2int8(_p6, _p7); + 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_kk0) + { + 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_kk0) + { + 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; + } + } + } + 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 = 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_kk0 = std::min(max_kk - 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; + 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); + + 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_kk0; kk += 4) + { + 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) + { + 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); + _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)); + p0a += 4; + p1a += 4; + p2a += 4; + p3a += 4; + if (psa) + psa += 4; + } + 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_kk0; kk++) + { + float v0 = *p0a++; + float v1 = *p1a++; + float v2 = *p2a++; + float v3 = *p3a++; + if (psa) + { + const float s = *psa++; + 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)); + } + + 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; + 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); + kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + v4f32 _p0 = (v4f32)__msa_ld_w(p0q, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1q, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2q, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3q, 0); + if (psq) + { + v4f32 _s = (v4f32)__msa_ld_w(psq, 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); + } + _p0 = __msa_fmul_w(_p0, _scale0); + _p1 = __msa_fmul_w(_p1, _scale1); + _p2 = __msa_fmul_w(_p2, _scale2); + _p3 = __msa_fmul_w(_p3, _scale3); + + ((int64_t*)pp)[0] = float2int8(_p0, _p1); + ((int64_t*)pp)[1] = float2int8(_p2, _p3); + pp += 16; + p0q += 4; + p1q += 4; + p2q += 4; + p3q += 4; + if (psq) + psq += 4; + } + if (kk + 1 < max_kk0) + { + 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; + v11 *= s1; + v20 *= s0; + v21 *= s1; + v30 *= s0; + v31 *= s1; + } + 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; + p0q += 2; + p1q += 2; + p2q += 2; + p3q += 2; + if (psq) + psq += 2; + kk += 2; + } + for (; kk < max_kk0; kk++) + { + float v0 = *p0q++; + float v1 = *p1q++; + float v2 = *p2q++; + float v3 = *p3q++; + if (psq) + { + const float s = *psq++; + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + } + 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 = 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_kk0 = std::min(max_kk - 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_kk0; kk++) + { + float v0 = *p0a++; + float v1 = *p1a++; + if (psa) + { + const float s = *psa++; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } + + 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_kk0; kk += 4) + { + 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_kk0) + { + float v00 = p0q[0]; + float v01 = p0q[1]; + float v10 = p1q[0]; + float v11 = p1q[1]; + if (psq) + { + 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_kk0; kk++) + { + float v0 = *p0q++; + float v1 = *p1q++; + if (psq) + { + const float s = *psq++; + v0 *= s; + v1 *= s; + } + 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 = A_data + (size_t)i0 * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk0 = std::min(max_kk - 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_kk0; kk++) + { + float v0 = *p0a++; + if (psa) + v0 *= *psa++; + absmax0 = std::max(absmax0, fabsf(v0)); + } + + 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_kk0; kk += 4) + { + float v0 = p0q[0]; + float v1 = p0q[1]; + float v2 = p0q[2]; + float v3 = p0q[3]; + if (psq) + { + 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_kk0) + { + float v0 = p0q[0]; + float v1 = p0q[1]; + if (psq) + { + 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_kk0; kk++) + { + float v0 = *p0q++; + if (psq) + v0 *= *psq++; + *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 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, k, max_kk, block_size, input_scale_ptr); + return; + } +#endif + + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + 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 + 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_kk0 = std::min(max_kk - 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_kk0; kk++) + { + v4f32 _p0 = (v4f32)__msa_ld_w(p0a, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p0a + 4, 0); + if (psa) + { + 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); + 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; + 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; + + const float* p0q = p0g; + const float* psq = sg; + 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); + v4f32 _scale4 = __msa_fill_w_f32(scale4); + v4f32 _scale5 = __msa_fill_w_f32(scale5); + v4f32 _scale6 = __msa_fill_w_f32(scale6); + v4f32 _scale7 = __msa_fill_w_f32(scale7); + kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + const float* p0 = p0q; + 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); + transpose4x4_ps(_p0, _p1, _p2, _p3); + + v4f32 _p4 = (v4f32)__msa_ld_w(p0 + 4, 0); + v4f32 _p5 = (v4f32)__msa_ld_w(p1 + 4, 0); + v4f32 _p6 = (v4f32)__msa_ld_w(p2 + 4, 0); + v4f32 _p7 = (v4f32)__msa_ld_w(p3 + 4, 0); + transpose4x4_ps(_p4, _p5, _p6, _p7); + if (psq) + { + v4f32 _s = (v4f32)__msa_ld_w(psq, 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); + } + _p0 = __msa_fmul_w(_p0, _scale0); + _p1 = __msa_fmul_w(_p1, _scale1); + _p2 = __msa_fmul_w(_p2, _scale2); + _p3 = __msa_fmul_w(_p3, _scale3); + _p4 = __msa_fmul_w(_p4, _scale4); + _p5 = __msa_fmul_w(_p5, _scale5); + _p6 = __msa_fmul_w(_p6, _scale6); + _p7 = __msa_fmul_w(_p7, _scale7); + + ((int64_t*)pp)[0] = float2int8(_p0, _p1); + ((int64_t*)pp)[1] = float2int8(_p2, _p3); + ((int64_t*)pp)[2] = float2int8(_p4, _p5); + ((int64_t*)pp)[3] = float2int8(_p6, _p7); + pp += 32; + p0q = p3 + A_hstep; + if (psq) + psq += 4; + } + if (kk + 1 < max_kk0) + { + 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; + 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)); + v4f32 _scale0123 = (v4f32)__msa_set_w(__msa_load_w(&scale0), __msa_load_w(&scale1), __msa_load_w(&scale2), __msa_load_w(&scale3)); + v4f32 _scale4567 = (v4f32)__msa_set_w(__msa_load_w(&scale4), __msa_load_w(&scale5), __msa_load_w(&scale6), __msa_load_w(&scale7)); + v16i8 _q0 = float2int8(__msa_fmul_w(_p0, _scale0123)); + v16i8 _q1 = float2int8(__msa_fmul_w(_p1, _scale0123)); + 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, _scale4567)); + _q1 = float2int8(__msa_fmul_w(_p1, _scale4567)); + 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; + p0q = p1 + A_hstep; + if (psq) + psq += 2; + kk += 2; + } + if (kk < max_kk0) + { + 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)); + v4f32 _scale0123 = (v4f32)__msa_set_w(__msa_load_w(&scale0), __msa_load_w(&scale1), __msa_load_w(&scale2), __msa_load_w(&scale3)); + v4f32 _scale4567 = (v4f32)__msa_set_w(__msa_load_w(&scale4), __msa_load_w(&scale5), __msa_load_w(&scale6), __msa_load_w(&scale7)); + v16i8 _q0 = float2int8(__msa_fmul_w(_p0, _scale0123)); + v16i8 _q1 = float2int8(__msa_fmul_w(_p1, _scale4567)); + ((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_kk0 = std::min(max_kk - 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_kk0; kk++) + { + 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); + 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; + + 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_kk0; kk += 4) + { + const float* p0 = p0q; + 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); + transpose4x4_ps(_p0, _p1, _p2, _p3); + if (psq) + { + v4f32 _s = (v4f32)__msa_ld_w(psq, 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); + } + _p0 = __msa_fmul_w(_p0, _scale0); + _p1 = __msa_fmul_w(_p1, _scale1); + _p2 = __msa_fmul_w(_p2, _scale2); + _p3 = __msa_fmul_w(_p3, _scale3); + + ((int64_t*)pp)[0] = float2int8(_p0, _p1); + ((int64_t*)pp)[1] = float2int8(_p2, _p3); + pp += 16; + p0q = p3 + A_hstep; + if (psq) + psq += 4; + } + if (kk + 1 < max_kk0) + { + const float* p0 = p0q; + 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 (psq) + { + const float s0 = psq[0]; + const float s1 = psq[1]; + v00 *= s0; + v10 *= s0; + v20 *= s0; + v30 *= s0; + v01 *= s1; + v11 *= s1; + v21 *= s1; + v31 *= s1; + } + 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; + p0q = p1 + A_hstep; + if (psq) + psq += 2; + kk += 2; + } + for (; kk < max_kk0; kk++) + { + float v0 = p0q[0]; + float v1 = p0q[1]; + float v2 = p0q[2]; + float v3 = p0q[3]; + if (psq) + { + const float s = *psq++; + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + } + 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; + } + } + } +#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_kk0 = std::min(max_kk - 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_kk0; kk++) + { + float v0 = p0a[0]; + float v1 = p0a[1]; + if (psa) + { + const float s = *psa++; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + p0a += A_hstep; + } + + 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_kk0; kk += 4) + { + float v00 = p0q[0]; + float v10 = p0q[1]; + float v01 = p0q[A_hstep]; + float v11 = p0q[A_hstep + 1]; + float v02 = p0q[A_hstep * 2]; + float v12 = p0q[A_hstep * 2 + 1]; + float v03 = p0q[A_hstep * 3]; + float v13 = p0q[A_hstep * 3 + 1]; + if (psq) + { + v00 *= psq[0]; + v10 *= psq[0]; + v01 *= psq[1]; + v11 *= psq[1]; + v02 *= psq[2]; + v12 *= psq[2]; + v03 *= psq[3]; + v13 *= psq[3]; + psq += 4; + } + 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); + p0q += A_hstep * 4; + pp += 8; + } + if (kk + 1 < max_kk0) + { + float v00 = p0q[0]; + float v10 = p0q[1]; + float v01 = p0q[A_hstep]; + float v11 = p0q[A_hstep + 1]; + if (psq) + { + v00 *= psq[0]; + v10 *= psq[0]; + v01 *= psq[1]; + v11 *= psq[1]; + psq += 2; + } + pp[0] = float2int8(v00 * scale0); + pp[1] = float2int8(v01 * scale0); + pp[2] = float2int8(v10 * scale1); + pp[3] = float2int8(v11 * scale1); + p0q += A_hstep * 2; + pp += 4; + kk += 2; + } + for (; kk < max_kk0; kk++) + { + float v0 = p0q[0]; + float v1 = p0q[1]; + if (psq) + { + const float s = *psq++; + v0 *= s; + v1 *= s; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + p0q += A_hstep; + } + } + } + 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_kk0 = std::min(max_kk - 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_kk0; kk++) + { + float v0 = *p0a; + if (psa) + v0 *= *psa++; + absmax0 = std::max(absmax0, fabsf(v0)); + p0a += A_hstep; + } + + 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_kk0; kk += 4) + { + float v0 = p0q[0]; + float v1 = p0q[A_hstep]; + float v2 = p0q[A_hstep * 2]; + float v3 = p0q[A_hstep * 3]; + if (psq) + { + v0 *= psq[0]; + v1 *= psq[1]; + v2 *= psq[2]; + v3 *= psq[3]; + psq += 4; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale0); + pp[2] = float2int8(v2 * scale0); + pp[3] = float2int8(v3 * scale0); + p0q += A_hstep * 4; + pp += 4; + } + if (kk + 1 < max_kk0) + { + float v0 = p0q[0]; + float v1 = p0q[A_hstep]; + if (psq) + { + v0 *= psq[0]; + v1 *= psq[1]; + psq += 2; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale0); + p0q += A_hstep * 2; + pp += 2; + kk += 2; + } + for (; kk < max_kk0; kk++) + { + 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 k, int max_kk, int K, 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, k, max_kk, K, block_size); + return; + } +#endif + + const signed char* pAT = AT_tile; + const int A_hstep = max_kk; + const float* pAT_descales = AT_descales_tile; + const int A_descales_hstep = (max_kk + block_size - 1) / block_size; + const signed char* pBT = BT_tile; + const float* pBT_descales = BT_descales_tile; + float* outptr = topT_tile; + const int block_count = (K + block_size - 1) / block_size; + const int block_start = k / block_size; + int ii = 0; +#if __mips_msa + for (; ii + 7 < max_ii; ii += 8) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + v8i16 _one = __msa_fill_h(1); + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + v4f32 _fsum1; + v4f32 _fsum2; + v4f32 _fsum3; + v4f32 _fsum4; + v4f32 _fsum5; + v4f32 _fsum6; + v4f32 _fsum7; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + _fsum2 = (v4f32)__msa_fill_w(0); + _fsum3 = (v4f32)__msa_fill_w(0); + _fsum4 = (v4f32)__msa_fill_w(0); + _fsum5 = (v4f32)__msa_fill_w(0); + _fsum6 = (v4f32)__msa_fill_w(0); + _fsum7 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum4 = (v4f32)__msa_ld_w(outptr + 4, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 8, 0); + _fsum5 = (v4f32)__msa_ld_w(outptr + 12, 0); + _fsum2 = (v4f32)__msa_ld_w(outptr + 16, 0); + _fsum6 = (v4f32)__msa_ld_w(outptr + 20, 0); + _fsum3 = (v4f32)__msa_ld_w(outptr + 24, 0); + _fsum7 = (v4f32)__msa_ld_w(outptr + 28, 0); + transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); + transpose4x4_ps(_fsum4, _fsum5, _fsum6, _fsum7); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; kk += 4) + { + __builtin_prefetch(pA + 64); + __builtin_prefetch(pB + 64); + 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); + + 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); + _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_kk0) + { + 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); + _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_kk0) + { + 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); + _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; + } + + v4f32 _descaleB = (v4f32)__msa_ld_w(pB_descales, 0); + v4f32 _descaleA0 = (v4f32)__msa_ld_w(pA_descales, 0); + v4f32 _descaleA1 = (v4f32)__msa_ld_w(pA_descales + 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); + pA_descales += 8; + pB_descales += 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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + v4f32 _fsum1; + v4f32 _fsum2; + v4f32 _fsum3; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + _fsum2 = (v4f32)__msa_fill_w(0); + _fsum3 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + _fsum2 = (v4f32)__msa_ld_w(outptr + 8, 0); + _fsum3 = (v4f32)__msa_ld_w(outptr + 12, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; kk += 4) + { + __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); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pA1, _pBr), _one); + pA += 32; + pB += 8; + } + + 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); + 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_kk0) + { + 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)); + _sum3x = __msa_addv_w(_sum3x, (v4i32)__msa_ilvl_h(_sign1, _s1)); + pA += 16; + pB += 4; + kk += 2; + } + if (kk < max_kk0) + { + 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)); + _sum3x = __msa_addv_w(_sum3x, (v4i32)__msa_ilvl_h(_sign1, _s1)); + pA += 8; + pB += 2; + } + + v4f32 _descaleA0 = (v4f32)__msa_ld_w(pA_descales, 0); + v4f32 _descaleA1 = (v4f32)__msa_ld_w(pA_descales + 4, 0); + v4f32 _scale = __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pB_descales[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0x), _scale); + _scale = __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pB_descales[0])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum2x), _scale); + _scale = __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pB_descales[1])); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum1x), _scale); + _scale = __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pB_descales[1])); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3x), _scale); + pA_descales += 8; + pB_descales += 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; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + v4f32 _fsum1; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __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; + pB += 4; + } + if (kk + 1 < max_kk0) + { + 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; + pB += 2; + kk += 2; + } + if (kk < max_kk0) + { + 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++; + } + + v4f32 _descaleA0 = (v4f32)__msa_ld_w(pA_descales, 0); + v4f32 _descaleA1 = (v4f32)__msa_ld_w(pA_descales + 4, 0); + v4f32 _scale = __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pB_descales[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pB_descales[0])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); + pA_descales += 8; + pB_descales++; + } + + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + pB_panel += K; + pB_descales_panel += block_count; + } + + 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_panel = pBT; + const float* pB_descales_panel = pBT_descales; + v8i16 _one = __msa_fill_h(1); + + int jj = 0; + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + v4f32 _fsum0; + v4f32 _fsum1; + v4f32 _fsum2; + v4f32 _fsum3; + v4f32 _fsum4; + v4f32 _fsum5; + v4f32 _fsum6; + v4f32 _fsum7; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + _fsum2 = (v4f32)__msa_fill_w(0); + _fsum3 = (v4f32)__msa_fill_w(0); + _fsum4 = (v4f32)__msa_fill_w(0); + _fsum5 = (v4f32)__msa_fill_w(0); + _fsum6 = (v4f32)__msa_fill_w(0); + _fsum7 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + _fsum2 = (v4f32)__msa_ld_w(outptr + 8, 0); + _fsum3 = (v4f32)__msa_ld_w(outptr + 12, 0); + _fsum4 = (v4f32)__msa_ld_w(outptr + 16, 0); + _fsum5 = (v4f32)__msa_ld_w(outptr + 20, 0); + _fsum6 = (v4f32)__msa_ld_w(outptr + 24, 0); + _fsum7 = (v4f32)__msa_ld_w(outptr + 28, 0); + transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); + transpose4x4_ps(_fsum4, _fsum5, _fsum6, _fsum7); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; kk += 4) + { + __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); + _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; + } + + _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_kk0) + { + 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); + _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)); + pA += 8; + pB0 += 8; + pB1 += 8; + kk += 2; + } + for (; kk < max_kk0; 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; + } + + v4f32 _descaleB0 = (v4f32)__msa_ld_w(pB_descales0, 0); + v4f32 _descaleB1 = (v4f32)__msa_ld_w(pB_descales1, 0); + v4f32 _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pA_descales[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pA_descales[1])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); + _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pA_descales[2])); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum2), _scale); + _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pA_descales[3])); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pA_descales[0])); + _fsum4 = __ncnn_msa_fmadd_w(_fsum4, (v4f32)__msa_ffint_s_w(_sum4), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pA_descales[1])); + _fsum5 = __ncnn_msa_fmadd_w(_fsum5, (v4f32)__msa_ffint_s_w(_sum5), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pA_descales[2])); + _fsum6 = __ncnn_msa_fmadd_w(_fsum6, (v4f32)__msa_ffint_s_w(_sum6), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pA_descales[3])); + _fsum7 = __ncnn_msa_fmadd_w(_fsum7, (v4f32)__msa_ffint_s_w(_sum7), _scale); + pA_descales += 4; + pB_descales0 += 4; + pB_descales1 += 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_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + v4f32 _fsum1; + v4f32 _fsum2; + v4f32 _fsum3; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + _fsum2 = (v4f32)__msa_fill_w(0); + _fsum3 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + _fsum2 = (v4f32)__msa_ld_w(outptr + 8, 0); + _fsum3 = (v4f32)__msa_ld_w(outptr + 12, 0); + transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; kk += 4) + { + __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); + _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_kk0) + { + 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); + _sum2_3 = __msa_dotp_s_h(_pA3, _pB); + pA += 8; + pB += 8; + kk += 2; + } + for (; kk < max_kk0; kk++) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); + 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); + 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)); + _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)); + v4f32 _descaleB = (v4f32)__msa_ld_w(pB_descales, 0); + v4f32 _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pA_descales[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pA_descales[1])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); + _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pA_descales[2])); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum2), _scale); + _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pA_descales[3])); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3), _scale); + pA_descales += 4; + pB_descales += 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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + v4f32 _fsum1; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __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; + pB += 8; + } + v8i16 _sum2_0 = __msa_fill_h(0); + v8i16 _sum2_1 = __msa_fill_h(0); + if (kk + 1 < max_kk0) + { + 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; + pB += 4; + kk += 2; + } + for (; kk < max_kk0; kk++) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); + 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; + } + 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)); + v4f32 _descaleA = (v4f32)__msa_ld_w(pA_descales, 0); + v4f32 _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pB_descales[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0x), _scale); + _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pB_descales[1])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1x), _scale); + pA_descales += 4; + pB_descales += 2; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum0 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __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; + } + v8i16 _sum2_0 = __msa_fill_h(0); + if (kk + 1 < max_kk0) + { + 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; + kk += 2; + } + for (; kk < max_kk0; kk++) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); + 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)); + v4f32 _descaleA = (v4f32)__msa_ld_w(pA_descales, 0); + v4f32 _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pB_descales[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + pA_descales += 4; + pB_descales++; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + outptr += 4; + pB_panel += K; + pB_descales_panel += block_count; + } + + 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_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __mips_msa + v8i16 _one = __msa_fill_h(1); + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + v4f32 _fsum0; + v4f32 _fsum1; + v4f32 _fsum2; + v4f32 _fsum3; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + _fsum2 = (v4f32)__msa_fill_w(0); + _fsum3 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + _fsum2 = (v4f32)__msa_ld_w(outptr + 8, 0); + _fsum3 = (v4f32)__msa_ld_w(outptr + 12, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; kk += 4) + { + __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); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pA, _pB11), _one); + pA += 8; + pB0 += 16; + pB1 += 16; + } + if (kk + 1 < max_kk0) + { + 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)); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvl_h(_sign2, _s2)); + pA += 4; + pB0 += 8; + pB1 += 8; + kk += 2; + } + if (kk < max_kk0) + { + 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)); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvl_h(_sign2, _s2)); + pA += 2; + pB0 += 4; + pB1 += 4; + } + v4f32 _descaleA = (v4f32)__msa_fill_d_ptr(pA_descales); + v4f32 _descaleB0 = (v4f32)__msa_set_w(__msa_load_w(pB_descales0), __msa_load_w(pB_descales0), __msa_load_w(pB_descales0 + 1), __msa_load_w(pB_descales0 + 1)); + v4f32 _descaleB1 = (v4f32)__msa_set_w(__msa_load_w(pB_descales0 + 2), __msa_load_w(pB_descales0 + 2), __msa_load_w(pB_descales0 + 3), __msa_load_w(pB_descales0 + 3)); + v4f32 _descaleB2 = (v4f32)__msa_set_w(__msa_load_w(pB_descales1), __msa_load_w(pB_descales1), __msa_load_w(pB_descales1 + 1), __msa_load_w(pB_descales1 + 1)); + v4f32 _descaleB3 = (v4f32)__msa_set_w(__msa_load_w(pB_descales1 + 2), __msa_load_w(pB_descales1 + 2), __msa_load_w(pB_descales1 + 3), __msa_load_w(pB_descales1 + 3)); + 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); + pA_descales += 2; + pB_descales0 += 4; + pB_descales1 += 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_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + v4f32 _fsum1; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + __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; + pB += 16; + } + if (kk + 1 < max_kk0) + { + 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; + pB += 8; + kk += 2; + } + if (kk < max_kk0) + { + 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; + } + v4f32 _descaleA = (v4f32)__msa_fill_d_ptr(pA_descales); + v4f32 _descaleB0 = (v4f32)__msa_set_w(__msa_load_w(pB_descales), __msa_load_w(pB_descales), __msa_load_w(pB_descales + 1), __msa_load_w(pB_descales + 1)); + v4f32 _descaleB1 = (v4f32)__msa_set_w(__msa_load_w(pB_descales + 2), __msa_load_w(pB_descales + 2), __msa_load_w(pB_descales + 3), __msa_load_w(pB_descales + 3)); + 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); + pA_descales += 2; + pB_descales += 4; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum; + if (k == 0) + { + _fsum = (v4f32)__msa_fill_w(0); + } + else + { + _fsum = (v4f32)__msa_ld_w(outptr, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + 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_kk0) + { + 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)__msa_fill_d_ptr(pA_descales); + v4f32 _descaleB = (v4f32)__msa_set_w(__msa_load_w(pB_descales), __msa_load_w(pB_descales), __msa_load_w(pB_descales + 1), __msa_load_w(pB_descales + 1)); + v4f32 _scale = __msa_fmul_w(_descaleA, _descaleB); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pA_descales += 2; + pB_descales += 2; + } + __msa_st_w((v4i32)_fsum, outptr, 0); + outptr += 4; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum; + if (k == 0) + { + _fsum = (v4f32)__msa_fill_w(0); + } + else + { + _fsum = (v4f32)__msa_loadl_d(outptr); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + 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_kk0) + { + 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)__msa_fill_d_ptr(pA_descales); + v4f32 _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pB_descales[0])); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pA_descales += 2; + pB_descales++; + } + __msa_storel_d((v4i32)_fsum, outptr); + outptr += 2; + pB_panel += K; + pB_descales_panel += block_count; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + float sum00; + float sum01; + float sum10; + float sum11; + if (k == 0) + { + sum00 = 0.f; + sum01 = 0.f; + sum10 = 0.f; + sum11 = 0.f; + } + else + { + sum00 = outptr[0]; + sum01 = outptr[1]; + sum10 = outptr[2]; + sum11 = outptr[3]; + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; 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_kk0; kk += 4) + { + __builtin_prefetch(pB + 32); + 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); + 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_kk0; 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_kk0) + { + 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_kk0; 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 * pA_descales[0] * pB_descales[0]; + sum01 += sum01_i * pA_descales[1] * pB_descales[0]; + sum10 += sum10_i * pA_descales[0] * pB_descales[1]; + sum11 += sum11_i * pA_descales[1] * pB_descales[1]; + pA_descales += 2; + pB_descales += 2; + } + + outptr[0] = sum00; + outptr[1] = sum01; + outptr[2] = sum10; + outptr[3] = sum11; + outptr += 4; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + float sum0; + float sum1; + if (k == 0) + { + sum0 = 0.f; + sum1 = 0.f; + } + else + { + sum0 = outptr[0]; + sum1 = outptr[1]; + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; 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_kk0; kk += 4) + { + __builtin_prefetch(pB + 16); + 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); + _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_kk0; 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_kk0) + { + 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_kk0; kk++) + { + sum0_i += pA[0] * pB[0]; + sum1_i += pA[1] * pB[0]; + pA += 2; + pB++; + } + sum0 += sum0_i * pA_descales[0] * pB_descales[0]; + sum1 += sum1_i * pA_descales[1] * pB_descales[0]; + pA_descales += 2; + pB_descales++; + } + + outptr[0] = sum0; + outptr[1] = sum1; + outptr += 2; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += (size_t)2 * A_hstep; + pAT_descales += (size_t)2 * A_descales_hstep; + } + for (; ii < max_ii; ii++) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __mips_msa + v8i16 _one = __msa_fill_h(1); + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + const signed char* pB0 = pB_panel + (size_t)4 * k; + const signed char* pB1 = pB_panel + (size_t)4 * K + (size_t)4 * k; + const float* pB_descales0 = pB_descales_panel + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales_panel + (size_t)4 * block_count + (size_t)4 * block_start; + v4f32 _fsum0; + v4f32 _fsum1; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + _fsum1 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + _fsum1 = (v4f32)__msa_ld_w(outptr + 4, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + 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_kk0; 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_kk0; kk += 4) + { + __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; + pB0 += 16; + pB1 += 16; + } + if (kk + 1 < max_kk0) + { + 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; + pB0 += 8; + pB1 += 8; + kk += 2; + } + if (kk < max_kk0) + { + 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; + } + v4f32 _descaleB0 = (v4f32)__msa_ld_w(pB_descales0, 0); + v4f32 _descaleB1 = (v4f32)__msa_ld_w(pB_descales1, 0); + v4f32 _descaleA = __msa_fill_w_f32(pA_descales[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); + pA_descales++; + pB_descales0 += 4; + pB_descales1 += 4; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum0; + if (k == 0) + { + _fsum0 = (v4f32)__msa_fill_w(0); + } + else + { + _fsum0 = (v4f32)__msa_ld_w(outptr, 0); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum0 = __msa_fill_w(0); + int kk = 0; + { + v4i32 _sum1 = __msa_fill_w(0); + for (; kk + 7 < max_kk0; 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_kk0; kk += 4) + { + __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_kk0) + { + 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; + kk += 2; + } + if (kk < max_kk0) + { + 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; + } + v4f32 _descaleB = (v4f32)__msa_ld_w(pB_descales, 0); + v4f32 _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pA_descales[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + pA_descales++; + pB_descales += 4; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + outptr += 4; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum; + if (k == 0) + { + _fsum = (v4f32)__msa_fill_w(0); + } + else + { + _fsum = (v4f32)__msa_loadl_d(outptr); + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + 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_kk0) + { + 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)__msa_fill_d_ptr(pB_descales); + v4f32 _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pA_descales[0])); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pA_descales++; + pB_descales += 2; + } + __msa_storel_d((v4i32)_fsum, outptr); + outptr += 2; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + v4f32 _fsum = (v4f32)__msa_fill_w(0); + if (k != 0) + _fsum[0] = outptr[0]; + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0) + { + 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_kk0) + { + 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(pA_descales[0] * pB_descales[0]); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pA_descales++; + pB_descales++; + } + *outptr++ = _fsum[0]; + pB_panel += K; + pB_descales_panel += block_count; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + float sum0; + float sum1; + if (k == 0) + { + sum0 = 0.f; + sum1 = 0.f; + } + else + { + sum0 = outptr[0]; + sum1 = outptr[1]; + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; 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_kk0; kk += 4) + { + __builtin_prefetch(pB + 32); + 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); + _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_kk0; 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_kk0) + { + 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_kk0; kk++) + { + sum0_i += pA[0] * pB[0]; + sum1_i += pA[0] * pB[1]; + pA++; + pB += 2; + } + sum0 += sum0_i * pA_descales[0] * pB_descales[0]; + sum1 += sum1_i * pA_descales[0] * pB_descales[1]; + pA_descales++; + pB_descales += 2; + } + + outptr[0] = sum0; + outptr[1] = sum1; + outptr += 2; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + float sum0; + if (k == 0) + { + sum0 = 0.f; + } + else + { + sum0 = outptr[0]; + } + + for (int kk0 = 0; kk0 < max_kk; kk0 += block_size) + { + const int max_kk0 = std::min(max_kk - kk0, 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_kk0; 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_kk0; kk += 4) + { + 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); + _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_kk0; 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_kk0) + { + sum0_i += pA[0] * pB[0] + pA[1] * pB[1]; + pA += 2; + pB += 2; + kk += 2; + } + for (; kk < max_kk0; kk++) + sum0_i += *pA++ * *pB++; + sum0 += sum0_i * pA_descales[0] * pB_descales[0]; + pA_descales++; + pB_descales++; + } + + *outptr++ = sum0; + pB_panel += K; + pB_descales_panel += block_count; + } + + 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, 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; + const float* pp = topT; + 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) + { + 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) + { + v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _c = (v4f32)__msa_ld_w(pC0, 0); + pC0 += 4; + _f0 = __msa_fadd_w(_f0, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC1, 0); + pC1 += 4; + _f1 = __msa_fadd_w(_f1, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC2, 0); + pC2 += 4; + _f2 = __msa_fadd_w(_f2, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC3, 0); + pC3 += 4; + _f3 = __msa_fadd_w(_f3, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC4, 0); + pC4 += 4; + _f4 = __msa_fadd_w(_f4, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC5, 0); + pC5 += 4; + _f5 = __msa_fadd_w(_f5, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC6, 0); + pC6 += 4; + _f6 = __msa_fadd_w(_f6, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC7, 0); + pC7 += 4; + _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); + pC += 4; + 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) + { + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + v4f32 _c4 = (v4f32)__msa_set_w(__msa_load_w(pC4), __msa_load_w(pC5), __msa_load_w(pC6), __msa_load_w(pC7)); + v4f32 _c1 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1), __msa_load_w(pC2 + 1), __msa_load_w(pC3 + 1)); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + v4f32 _c5 = (v4f32)__msa_set_w(__msa_load_w(pC4 + 1), __msa_load_w(pC5 + 1), __msa_load_w(pC6 + 1), __msa_load_w(pC7 + 1)); + pC4 += 2; + pC5 += 2; + pC6 += 2; + pC7 += 2; + if (beta != 1.f) + { + 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]; + pC += 2; + 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) + { + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + pC0++; + pC1++; + pC2++; + pC3++; + v4f32 _c4 = (v4f32)__msa_set_w(__msa_load_w(pC4), __msa_load_w(pC5), __msa_load_w(pC6), __msa_load_w(pC7)); + pC4++; + pC5++; + pC6++; + pC7++; + 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, _c0); + _f4 = __msa_fadd_w(_f4, _c4); + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + pC++; + 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) + { + 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); + pC0 += 8; + v4f32 _c5 = (v4f32)__msa_ld_w(pC1 + 4, 0); + pC1 += 8; + v4f32 _c6 = (v4f32)__msa_ld_w(pC2 + 4, 0); + pC2 += 8; + v4f32 _c7 = (v4f32)__msa_ld_w(pC3 + 4, 0); + pC3 += 8; + 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); + pC += 8; + 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; + } + 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); + pC0 += 4; + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + pC1 += 4; + v4f32 _c2 = (v4f32)__msa_ld_w(pC2, 0); + pC2 += 4; + v4f32 _c3 = (v4f32)__msa_ld_w(pC3, 0); + pC3 += 4; + 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); + pC += 4; + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + v4f32 _c1 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1), __msa_load_w(pC2 + 1), __msa_load_w(pC3 + 1)); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + 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]; + pC += 2; + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + pC0++; + pC1++; + pC2++; + pC3++; + 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]; + pC++; + 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++; + } + 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); + pC0 += 8; + v4f32 _c2 = (v4f32)__msa_ld_w(pC1, 0); + v4f32 _c3 = (v4f32)__msa_ld_w(pC1 + 4, 0); + pC1 += 8; + 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); + pC += 8; + 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; + } + 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); + pC0 += 4; + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + pC1 += 4; + 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); + pC += 4; + 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; + } + 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)__msa_set_w(__msa_load_w(&c0), __msa_load_w(&c1), __msa_load_w(&c0), __msa_load_w(&c1))); + if (broadcast_type_C == 3) + { + v4f32 _c = (v4f32)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1)); + pC0 += 2; + pC1 += 2; + 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]; + pC += 2; + if (beta != 1.f) + { + cc0 *= beta; + cc1 *= beta; + } + _f = __msa_fadd_w(_f, (v4f32)__msa_set_w(__msa_load_w(&cc0), __msa_load_w(&cc0), __msa_load_w(&cc1), __msa_load_w(&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; + } +#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]; + pC0 += 2; + float c11 = pC1[1]; + pC1 += 2; + 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]; + pC += 2; + 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; + } + 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]; + pC0++; + float c1 = pC1[0]; + pC1++; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + pC++; + 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++; + } + 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); + pC0 += 8; + 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; + } + 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); + pC0 += 4; + 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; + } + 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]); + pC0 += 2; + 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; + } +#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]; + pC0 += 2; + 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]; + pC0 += 2; + 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; + } + 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]; + pC0++; + } + 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++; + } + outptr += out_hstep; + } +} + +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, 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; + const float* pp = topT; + 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) + { + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + v4f32 _ch0 = (v4f32)__msa_set_w(__msa_load_w(pC4), __msa_load_w(pC5), __msa_load_w(pC6), __msa_load_w(pC7)); + v4f32 _cl1 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1), __msa_load_w(pC2 + 1), __msa_load_w(pC3 + 1)); + v4f32 _ch1 = (v4f32)__msa_set_w(__msa_load_w(pC4 + 1), __msa_load_w(pC5 + 1), __msa_load_w(pC6 + 1), __msa_load_w(pC7 + 1)); + v4f32 _cl2 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 2), __msa_load_w(pC1 + 2), __msa_load_w(pC2 + 2), __msa_load_w(pC3 + 2)); + v4f32 _ch2 = (v4f32)__msa_set_w(__msa_load_w(pC4 + 2), __msa_load_w(pC5 + 2), __msa_load_w(pC6 + 2), __msa_load_w(pC7 + 2)); + v4f32 _cl3 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 3), __msa_load_w(pC1 + 3), __msa_load_w(pC2 + 3), __msa_load_w(pC3 + 3)); + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; + v4f32 _ch3 = (v4f32)__msa_set_w(__msa_load_w(pC4 + 3), __msa_load_w(pC5 + 3), __msa_load_w(pC6 + 3), __msa_load_w(pC7 + 3)); + pC4 += 4; + pC5 += 4; + pC6 += 4; + pC7 += 4; + if (beta != 1.f) + { + 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); + pC += 4; + 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) + { + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + v4f32 _c4 = (v4f32)__msa_set_w(__msa_load_w(pC4), __msa_load_w(pC5), __msa_load_w(pC6), __msa_load_w(pC7)); + v4f32 _c1 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1), __msa_load_w(pC2 + 1), __msa_load_w(pC3 + 1)); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + v4f32 _c5 = (v4f32)__msa_set_w(__msa_load_w(pC4 + 1), __msa_load_w(pC5 + 1), __msa_load_w(pC6 + 1), __msa_load_w(pC7 + 1)); + pC4 += 2; + pC5 += 2; + pC6 += 2; + pC7 += 2; + if (beta != 1.f) + { + 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]; + pC += 2; + 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) + { + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + pC0++; + pC1++; + pC2++; + pC3++; + v4f32 _c4 = (v4f32)__msa_set_w(__msa_load_w(pC4), __msa_load_w(pC5), __msa_load_w(pC6), __msa_load_w(pC7)); + pC4++; + pC5++; + pC6++; + pC7++; + 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, _c0); + _f4 = __msa_fadd_w(_f4, _c4); + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + pC++; + 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) + { + 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); + pC0 += 8; + v4f32 _c5 = (v4f32)__msa_ld_w(pC1 + 4, 0); + pC1 += 8; + v4f32 _c6 = (v4f32)__msa_ld_w(pC2 + 4, 0); + pC2 += 8; + v4f32 _c7 = (v4f32)__msa_ld_w(pC3 + 4, 0); + pC3 += 8; + 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); + pC += 8; + 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; + } + 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); + pC0 += 4; + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + pC1 += 4; + v4f32 _c2 = (v4f32)__msa_ld_w(pC2, 0); + pC2 += 4; + v4f32 _c3 = (v4f32)__msa_ld_w(pC3, 0); + pC3 += 4; + 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); + pC += 4; + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + v4f32 _c1 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1), __msa_load_w(pC2 + 1), __msa_load_w(pC3 + 1)); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + 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]; + pC += 2; + 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; + } + 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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC2), __msa_load_w(pC3)); + pC0++; + pC1++; + pC2++; + pC3++; + 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]; + pC++; + 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; + } + 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)__msa_set_w(__msa_load_w(&c0), __msa_load_w(&c1), __msa_load_w(&c0), __msa_load_w(&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)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1)); + v4f32 _c1 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 2), __msa_load_w(pC1 + 2), __msa_load_w(pC0 + 3), __msa_load_w(pC1 + 3)); + v4f32 _c2 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 4), __msa_load_w(pC1 + 4), __msa_load_w(pC0 + 5), __msa_load_w(pC1 + 5)); + v4f32 _c3 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 6), __msa_load_w(pC1 + 6), __msa_load_w(pC0 + 7), __msa_load_w(pC1 + 7)); + pC0 += 8; + pC1 += 8; + 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]; + pC += 8; + 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)__msa_set_w(__msa_load_w(&c00), __msa_load_w(&c00), __msa_load_w(&c01), __msa_load_w(&c01))); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_set_w(__msa_load_w(&c02), __msa_load_w(&c02), __msa_load_w(&c03), __msa_load_w(&c03))); + _f2 = __msa_fadd_w(_f2, (v4f32)__msa_set_w(__msa_load_w(&c04), __msa_load_w(&c04), __msa_load_w(&c05), __msa_load_w(&c05))); + _f3 = __msa_fadd_w(_f3, (v4f32)__msa_set_w(__msa_load_w(&c06), __msa_load_w(&c06), __msa_load_w(&c07), __msa_load_w(&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; + } + 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)__msa_set_w(__msa_load_w(&c0), __msa_load_w(&c1), __msa_load_w(&c0), __msa_load_w(&c1)); + _f0 = __msa_fadd_w(_f0, _c); + _f1 = __msa_fadd_w(_f1, _c); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1)); + v4f32 _c1 = (v4f32)__msa_set_w(__msa_load_w(pC0 + 2), __msa_load_w(pC1 + 2), __msa_load_w(pC0 + 3), __msa_load_w(pC1 + 3)); + pC0 += 4; + pC1 += 4; + 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]; + pC += 4; + if (beta != 1.f) + { + c00 *= beta; + c01 *= beta; + c02 *= beta; + c03 *= beta; + } + _f0 = __msa_fadd_w(_f0, (v4f32)__msa_set_w(__msa_load_w(&c00), __msa_load_w(&c00), __msa_load_w(&c01), __msa_load_w(&c01))); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_set_w(__msa_load_w(&c02), __msa_load_w(&c02), __msa_load_w(&c03), __msa_load_w(&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; + } + 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)__msa_set_w(__msa_load_w(&c0), __msa_load_w(&c1), __msa_load_w(&c0), __msa_load_w(&c1))); + if (broadcast_type_C == 3) + { + v4f32 _c = (v4f32)__msa_set_w(__msa_load_w(pC0), __msa_load_w(pC1), __msa_load_w(pC0 + 1), __msa_load_w(pC1 + 1)); + pC0 += 2; + pC1 += 2; + 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]; + pC += 2; + if (beta != 1.f) + { + cc0 *= beta; + cc1 *= beta; + } + _f = __msa_fadd_w(_f, (v4f32)__msa_set_w(__msa_load_w(&cc0), __msa_load_w(&cc0), __msa_load_w(&cc1), __msa_load_w(&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; + } +#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]; + pC0 += 2; + float c11 = pC1[1]; + pC1 += 2; + 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]; + pC += 2; + 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; + } + 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]; + pC0++; + float c1 = pC1[0]; + pC1++; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + pC++; + 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; + } + 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); + pC0 += 8; + 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; + } + 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); + pC0 += 4; + 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; + } + 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]); + pC0 += 2; + 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; + } +#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]; + pC0 += 2; + 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; + } + 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]; + pC0++; + if (beta != 1.f) + c *= beta; + } + sum0 += c; + } + if (alpha != 1.f) + sum0 *= alpha; + outptr[0] = sum0; + outptr += out_hstep; + } + outptr0++; + } +} + +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) +{ + const int l2_cache_size_int8 = (int)(get_cpu_level2_cache_size() / sizeof(signed char)); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + +#if __mips_msa + const int tile_m_align = 8; + const int tile_n_align = 8; +#else + const int tile_m_align = 2; + const int tile_n_align = 2; +#endif + + { +#if __mips_msa + int tile_size = (l2_cache_size_int8 - 16) / 8; +#else + int tile_size = (l2_cache_size_int8 - 2) / 3; +#endif + TILE_K = std::max(block_size, tile_size / block_size * block_size); + + if (K > 0) + { + int nn_K = (K + TILE_K - 1) / TILE_K; + TILE_K = std::min(TILE_K, ((K + nn_K - 1) / nn_K + block_size - 1) / block_size * block_size); + if (TILE_K >= K) + TILE_K = K; + } + } + + { + int tile_size = (l2_cache_size_int8 - tile_n_align * TILE_K) / std::max(1, TILE_K + tile_n_align); + TILE_M = std::max(tile_m_align, tile_size / tile_m_align * tile_m_align); + + if (M > 0) + { + int nn_M = std::max(std::min(nT, get_physical_cpu_count()), (M + TILE_M - 1) / TILE_M); + TILE_M = std::max(tile_m_align, std::min(TILE_M, ((M + nn_M - 1) / nn_M + tile_m_align - 1) / tile_m_align * tile_m_align)); + } + } + + if (N > 0) + { + int tile_size = TILE_K >= K ? (l2_cache_size_int8 - TILE_M * TILE_K) / std::max(1, TILE_K) : (l2_cache_size_int8 - TILE_M * TILE_K) / std::max(1, TILE_M + TILE_K); + TILE_N = std::max(tile_n_align, tile_size / tile_n_align * tile_n_align); + int nn_N = (N + TILE_N - 1) / TILE_N; + TILE_N = std::max(tile_n_align, std::min(TILE_N, ((N + nn_N - 1) / nn_N + tile_n_align - 1) / tile_n_align * tile_n_align)); + } + else + TILE_N = tile_n_align; + + if (constant_TILE_M > 0) + TILE_M = (constant_TILE_M + tile_m_align - 1) / tile_m_align * tile_m_align; + 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 - 1) / block_size * block_size); + if (K > 0 && TILE_K >= K) + TILE_K = K; + } +} diff --git a/src/layer/mips/multiheadattention_mips.cpp b/src/layer/mips/multiheadattention_mips.cpp index 16453933fd6..e07c82d5858 100644 --- a/src/layer/mips/multiheadattention_mips.cpp +++ b/src/layer/mips/multiheadattention_mips.cpp @@ -28,10 +28,362 @@ 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; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -260,13 +612,33 @@ int MultiHeadAttention_mips::create_pipeline(const Option& _opt) int MultiHeadAttention_mips::destroy_pipeline(const Option& _opt) { if (weight_block_quantize) - return 0; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 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 +649,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; } @@ -322,7 +694,17 @@ 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) - return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + } int q_blob_i = 0; int k_blob_i = 0; @@ -340,10 +722,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 +780,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 +790,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 +818,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 +869,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 +897,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 +944,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 bdb1bbbeab9..3bdaf3830ed 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 47663cee217..c7a1dccd3ac 100644 --- a/src/layer/multiheadattention.cpp +++ b/src/layer/multiheadattention.cpp @@ -8,41 +8,28 @@ namespace ncnn { -static bool mha_is_weight_block_quantize(int quantize_term) +int MultiHeadAttention::get_weight_block_quantize_params(int& weight_bits, int& block_size, bool& has_input_scale) const { - const int weight_bits = quantize_term / 100; + weight_bits = quantize_term / 100; const int format_code = quantize_term % 100 / 10; const int block_size_code = quantize_term % 10; if (weight_bits != 4 && weight_bits != 6 && weight_bits != 8) - return false; + return -1; if (format_code != 0 && format_code != 1) - return false; + return -1; if (block_size_code < 0 || block_size_code > 2) - return false; + return -1; - return true; -} + block_size = block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + has_input_scale = format_code == 1; -#if NCNN_WEIGHT_QUANT -static bool mha_weight_quantize_has_input_scale(int quantize_term) -{ - return quantize_term % 100 / 10 == 1; -} - -static int mha_weight_quantize_bits(int quantize_term) -{ - return quantize_term / 100; -} - -static int mha_weight_quantize_block_size(int quantize_term) -{ - const int block_size_code = quantize_term % 10; - return block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + return 0; } +#if NCNN_WEIGHT_QUANT static int mha_weight_quantize_packed_k_bytes(int constantK, int weight_bits) { if (constantK <= 0 || weight_bits <= 0) @@ -72,7 +59,10 @@ int MultiHeadAttention::load_param(const ParamDict& pd) scale = pd.get(6, 1.f / sqrtf(embed_dim / num_heads)); kv_cache = pd.get(7, 0); quantize_term = pd.get(18, 0); - weight_block_quantize = mha_is_weight_block_quantize(quantize_term); + int weight_bits; + int block_size; + bool has_input_scale; + weight_block_quantize = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) == 0; if (quantize_term == 4 || quantize_term == 5 || quantize_term == 6) { @@ -121,11 +111,14 @@ int MultiHeadAttention::load_model(const ModelBin& mb) const int qdim = weight_data_size / embed_dim; #if NCNN_WEIGHT_QUANT + int weight_bits = 0; + int block_size = 0; + bool has_input_scale = false; + if (weight_block_quantize && get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + if (weight_block_quantize) { - const int weight_bits = mha_weight_quantize_bits(quantize_term); - const int block_size = mha_weight_quantize_block_size(quantize_term); - const int q_packed_k_bytes = mha_weight_quantize_packed_k_bytes(qdim, weight_bits); const int k_packed_k_bytes = mha_weight_quantize_packed_k_bytes(kdim, weight_bits); const int v_packed_k_bytes = mha_weight_quantize_packed_k_bytes(vdim, weight_bits); @@ -172,7 +165,7 @@ int MultiHeadAttention::load_model(const ModelBin& mb) if (q_weight_data_quantize_scales.empty() || k_weight_data_quantize_scales.empty() || v_weight_data_quantize_scales.empty() || out_weight_data_quantize_scales.empty()) return -100; - if (mha_weight_quantize_has_input_scale(quantize_term)) + if (has_input_scale) { q_weight_data_input_scales = mb.load(qdim, 1); k_weight_data_input_scales = mb.load(kdim, 1); @@ -558,6 +551,122 @@ 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; + } + + const float scale = 127.f / absmax; + 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]; + 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; @@ -586,9 +695,11 @@ int MultiHeadAttention::forward_weight_block_quantize(const std::vector& bo const int embed_dim_per_head = embed_dim / num_heads; const int qdim = weight_data_size / embed_dim; - const int weight_bits = mha_weight_quantize_bits(quantize_term); - const int block_size = mha_weight_quantize_block_size(quantize_term); - const bool has_input_scale = mha_weight_quantize_has_input_scale(quantize_term); + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; const float* q_input_scale_ptr = has_input_scale ? (const float*)q_weight_data_input_scales : 0; const float* k_input_scale_ptr = has_input_scale ? (const float*)k_weight_data_input_scales : 0; @@ -603,35 +714,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 +777,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 +840,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 +1004,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/multiheadattention.h b/src/layer/multiheadattention.h index 6298415c5e7..893fcf5ce9a 100644 --- a/src/layer/multiheadattention.h +++ b/src/layer/multiheadattention.h @@ -22,6 +22,8 @@ class MultiHeadAttention : public Layer protected: void resolve_bottom_blob_index(int bottom_blob_count, int& q_blob_i, int& k_blob_i, int& v_blob_i, int& attn_mask_i, int& cached_xk_i, int& cached_xv_i) const; + int get_weight_block_quantize_params(int& weight_bits, int& block_size, bool& has_input_scale) const; + #if NCNN_WEIGHT_QUANT int forward_weight_block_quantize(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; #endif diff --git a/src/layer/riscv/gemm_riscv.cpp b/src/layer/riscv/gemm_riscv.cpp index 9e03af87279..8756195e920 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,10 +1873,303 @@ 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, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + 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); + + 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; + + const int nn_MK = nn_M * nn_K; + #pragma omp parallel for num_threads(nT) + 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_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, 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); + } + + 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 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, 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); + 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 topT_tile = topT.channel(get_omp_thread_num()); + 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); + + 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); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + } + } + } + + return 0; +} + +int Gemm_riscv::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; + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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, block_size, opt); + if (ret != 0) + return ret; + if (B_data_packed.empty() || B_data_descales.empty()) + return -100; + + BT_data_wq_int8 = B_data_packed; + BT_data_wq_int8_descales = B_data_descales; + + B_data.release(); + B_data_quantize_scales.release(); + + return 0; +} + +int Gemm_riscv::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; + } + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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) + { + if (output_N1M) + top_blob.create(M, 1, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + } + else + { + if (output_N1M) + top_blob.create(N, 1, M, (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, 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_riscv::create_pipeline(const Option& opt) { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + if (weight_bits == 8) + return create_pipeline_wq_int8(opt); +#endif // NCNN_WEIGHT_QUANT + return 0; } @@ -2019,10 +2316,30 @@ int Gemm_riscv::create_pipeline(const Option& opt) return 0; } +int Gemm_riscv::destroy_pipeline(const Option& /*opt*/) +{ +#if NCNN_WEIGHT_QUANT + BT_data_wq_int8.release(); + BT_data_wq_int8_descales.release(); +#endif + + return 0; +} + 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 + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + if (weight_bits == 8) + return forward_wq_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 2ef61f26892..01266544b55 100644 --- a/src/layer/riscv/gemm_riscv.h +++ b/src/layer/riscv/gemm_riscv.h @@ -14,10 +14,15 @@ 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 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_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 +33,10 @@ class Gemm_riscv : 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 }; } // 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 00000000000..1db74dca875 --- /dev/null +++ b/src/layer/riscv/gemm_wq_int8.h @@ -0,0 +1,2944 @@ +// 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) +{ + const size_t B_hstep = B.w; +#if __riscv_vector + const int packn = csrr_vlenb(); + const ptrdiff_t B_stride = (ptrdiff_t)B_hstep; + const size_t vl4 = __riscv_vsetvl_e8m1(4); + const size_t vl2 = __riscv_vsetvl_e8m1(2); +#else + const int packn = 4; +#endif // __riscv_vector + 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) + { + const signed char* p0 = B.row(j + jj); + 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); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { +#if __riscv_vector + __riscv_vse8_v_i8m1(pp, __riscv_vlse8_v_i8m1(p0, B_stride, vl4), vl4); + __riscv_vse8_v_i8m1(pp + 4, __riscv_vlse8_v_i8m1(p0 + 1, B_stride, vl4), vl4); + __riscv_vse8_v_i8m1(pp + 8, __riscv_vlse8_v_i8m1(p0 + 2, B_stride, vl4), vl4); + __riscv_vse8_v_i8m1(pp + 12, __riscv_vlse8_v_i8m1(p0 + 3, B_stride, vl4), vl4); +#else + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p0[B_hstep]; + pp[5] = p0[B_hstep + 1]; + pp[6] = p0[B_hstep + 2]; + pp[7] = p0[B_hstep + 3]; + pp[8] = p0[B_hstep * 2]; + pp[9] = p0[B_hstep * 2 + 1]; + pp[10] = p0[B_hstep * 2 + 2]; + pp[11] = p0[B_hstep * 2 + 3]; + pp[12] = p0[B_hstep * 3]; + pp[13] = p0[B_hstep * 3 + 1]; + pp[14] = p0[B_hstep * 3 + 2]; + pp[15] = p0[B_hstep * 3 + 3]; +#endif // __riscv_vector + p0 += 4; + pp += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { +#if __riscv_vector + __riscv_vse8_v_i8m1(pp, __riscv_vlse8_v_i8m1(p0, B_stride, vl4), vl4); + __riscv_vse8_v_i8m1(pp + 4, __riscv_vlse8_v_i8m1(p0 + 1, B_stride, vl4), vl4); +#else + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[B_hstep]; + pp[3] = p0[B_hstep + 1]; + pp[4] = p0[B_hstep * 2]; + pp[5] = p0[B_hstep * 2 + 1]; + pp[6] = p0[B_hstep * 3]; + pp[7] = p0[B_hstep * 3 + 1]; +#endif // __riscv_vector + p0 += 2; + pp += 8; + } + for (; kk < max_kk; kk++) + { +#if __riscv_vector + __riscv_vse8_v_i8m1(pp, __riscv_vlse8_v_i8m1(p0, B_stride, vl4), vl4); +#else + pp[0] = p0[0]; + pp[1] = p0[B_hstep]; + pp[2] = p0[B_hstep * 2]; + pp[3] = p0[B_hstep * 3]; +#endif // __riscv_vector + p0++; + pp += 4; + } + + 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 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); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { +#if __riscv_vector + __riscv_vse8_v_i8m1(pp, __riscv_vlse8_v_i8m1(p0, B_stride, vl2), vl2); + __riscv_vse8_v_i8m1(pp + 2, __riscv_vlse8_v_i8m1(p0 + 1, B_stride, vl2), vl2); + __riscv_vse8_v_i8m1(pp + 4, __riscv_vlse8_v_i8m1(p0 + 2, B_stride, vl2), vl2); + __riscv_vse8_v_i8m1(pp + 6, __riscv_vlse8_v_i8m1(p0 + 3, B_stride, vl2), vl2); +#else + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p0[B_hstep]; + pp[5] = p0[B_hstep + 1]; + pp[6] = p0[B_hstep + 2]; + pp[7] = p0[B_hstep + 3]; +#endif // __riscv_vector + p0 += 4; + pp += 8; + } + for (; kk + 1 < max_kk; kk += 2) + { +#if __riscv_vector + __riscv_vse8_v_i8m1(pp, __riscv_vlse8_v_i8m1(p0, B_stride, vl2), vl2); + __riscv_vse8_v_i8m1(pp + 2, __riscv_vlse8_v_i8m1(p0 + 1, B_stride, vl2), vl2); +#else + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[B_hstep]; + pp[3] = p0[B_hstep + 1]; +#endif // __riscv_vector + p0 += 2; + pp += 4; + } + for (; kk < max_kk; kk++) + { +#if __riscv_vector + __riscv_vse8_v_i8m1(pp, __riscv_vlse8_v_i8m1(p0, B_stride, vl2), vl2); +#else + pp[0] = p0[0]; + pp[1] = p0[B_hstep]; +#endif // __riscv_vector + p0++; + pp += 2; + } + + 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 max_kk = std::min(K - g * block_size, block_size); + for (int kk = 0; kk < max_kk; kk++) + *pp++ = *p0++; + *pd++ = 1.f / *ps0++; + } + } + } + + 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 k, int max_kk, int block_size, const float* input_scale_ptr) +{ + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + 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 + 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* p0g = A_data + (size_t)(i + ii) * A_hstep; + const float* psg = input_scale_ptr; + + for (int g = 0; g < block_count; g++) + { + const int block_kk = std::min(max_kk - g * block_size, block_size); + vfloat32m1_t _absmax = __riscv_vfmv_v_f_f32m1(0.f, vl); + + for (int kk = 0; kk < block_kk; kk++) + { + 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); + } + + 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 < block_kk; kk++) + { + 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); + vint32m4_t _v32x4 = __riscv_vundefined_i32m4(); + _v32x4 = __riscv_vset_v_i32m1_i32m4(_v32x4, 0, _v32); + vint16m2_t _v16 = __riscv_vnclip_wx_i16m2(_v32x4, 0, __RISCV_VXRM_RNU, vl); + __riscv_vse8_v_i8m1(pp, __riscv_vnclip_wx_i8m1(_v16, 0, __RISCV_VXRM_RNU, vl), vl); + pp += packn; + } + p0g += block_kk; + if (psg) + psg += block_kk; + } + } +#endif // __riscv_vector + for (; ii + 1 < max_ii; ii += 2) + { + const float* p0g = A_data + (size_t)(i + ii) * A_hstep; + const float* p1g = p0g + A_hstep; + const float* psg = input_scale_ptr; + + for (int g = 0; g < block_count; g++) + { + const int block_kk = std::min(max_kk - g * block_size, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + + int kk = 0; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0g + kk, vl); + vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1g + kk, vl); + if (psg) + { + 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); + } + _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; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v0 = p0g[kk]; + float v1 = p1g[kk]; + if (psg) + { + v0 *= psg[kk]; + v1 *= psg[kk]; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } + + 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; + + kk = 0; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0g + kk, vl); + vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1g + kk, vl); + if (psg) + { + 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); + } + vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); + vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); + vint8m2x2_t _q = __riscv_vcreate_v_i8m2x2(_q0, _q1); + __riscv_vsseg2e8_v_i8m2x2(pp, _q, vl); + pp += vl * 2; + kk += vl; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v0 = p0g[kk]; + float v1 = p1g[kk]; + if (psg) + { + v0 *= psg[kk]; + v1 *= psg[kk]; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } + p0g += block_kk; + p1g += block_kk; + if (psg) + psg += block_kk; + } + } + for (; ii < max_ii; ii++) + { + const float* p0g = A_data + (size_t)(i + ii) * A_hstep; + const float* psg = input_scale_ptr; + + for (int g = 0; g < block_count; g++) + { + const int block_kk = std::min(max_kk - g * block_size, block_size); + float absmax = 0.f; + + int kk = 0; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + 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; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v = p0g[kk]; + if (psg) + v *= psg[kk]; + absmax = std::max(absmax, fabsf(v)); + } + + const float scale = absmax == 0.f ? 0.f : 127.f / absmax; + *pd++ = absmax / 127.f; + + kk = 0; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + 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; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v = p0g[kk]; + if (psg) + v *= psg[kk]; + *pp++ = float2int8(v * scale); + } + p0g += block_kk; + if (psg) + psg += block_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 k, int max_kk, int block_size, const float* input_scale_ptr) +{ + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + 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 + const int packn = csrr_vlenb() / 4; + const size_t vl = __riscv_vsetvl_e32m1(packn); + for (; ii + (packn - 1) < max_ii; ii += packn) + { + const float* p0g = A_data + i + ii; + const float* psg = input_scale_ptr; + + for (int g = 0; g < block_count; g++) + { + const int block_kk = std::min(max_kk - 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 < block_kk; kk++) + { + 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 < block_kk; kk++) + { + 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); + vint32m4_t _v32x4 = __riscv_vundefined_i32m4(); + _v32x4 = __riscv_vset_v_i32m1_i32m4(_v32x4, 0, _v32); + vint16m2_t _v16 = __riscv_vnclip_wx_i16m2(_v32x4, 0, __RISCV_VXRM_RNU, vl); + __riscv_vse8_v_i8m1(pp, __riscv_vnclip_wx_i8m1(_v16, 0, __RISCV_VXRM_RNU, vl), vl); + pp += packn; + pAk += A_hstep; + } + p0g += (size_t)block_kk * A_hstep; + if (psg) + psg += block_kk; + } + } +#endif // __riscv_vector + for (; ii + 1 < max_ii; ii += 2) + { + const float* p0g = A_data + i + ii; + const float* psg = input_scale_ptr; + for (int g = 0; g < block_count; g++) + { + const int block_kk = std::min(max_kk - g * block_size, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + + int kk = 0; + const float* pAk = p0g; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + 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) + { + 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); + } + _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)); + pAk += vl * A_hstep; + kk += vl; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v0 = pAk[0]; + float v1 = pAk[1]; + if (psg) + { + const float s = psg[kk]; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + pAk += A_hstep; + } + + 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; + + kk = 0; + pAk = p0g; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + 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) + { + 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); + } + vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); + vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); + vint8m2x2_t _q = __riscv_vcreate_v_i8m2x2(_q0, _q1); + __riscv_vsseg2e8_v_i8m2x2(pp, _q, vl); + pp += vl * 2; + pAk += vl * A_hstep; + kk += vl; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v0 = pAk[0]; + float v1 = pAk[1]; + if (psg) + { + const float s = psg[kk]; + v0 *= s; + v1 *= s; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + pAk += A_hstep; + } + p0g += (size_t)block_kk * A_hstep; + if (psg) + psg += block_kk; + } + } + for (; ii < max_ii; ii++) + { + const float* p0g = A_data + i + ii; + const float* psg = input_scale_ptr; + + for (int g = 0; g < block_count; g++) + { + const int block_kk = std::min(max_kk - g * block_size, block_size); + float absmax = 0.f; + + int kk = 0; + const float* pAk = p0g; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + 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; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v = *pAk; + if (psg) + v *= psg[kk]; + absmax = std::max(absmax, fabsf(v)); + pAk += A_hstep; + } + + const float scale = absmax == 0.f ? 0.f : 127.f / absmax; + *pd++ = absmax / 127.f; + + kk = 0; + pAk = p0g; +#if __riscv_vector + while (kk < block_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(block_kk - kk); + 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; + } +#endif // __riscv_vector + for (; kk < block_kk; kk++) + { + float v = *pAk; + if (psg) + v *= psg[kk]; + *pp++ = float2int8(v * scale); + pAk += A_hstep; + } + p0g += (size_t)block_kk * A_hstep; + if (psg) + psg += block_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 k, int max_kk, int K, 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 A_hstep = max_kk; + const int A_descales_hstep = (max_kk + block_size - 1) / block_size; + const int block_count = (K + block_size - 1) / block_size; + const int block_start = k / block_size; + + int ii = 0; +#if __riscv_vector + const int packn = csrr_vlenb() / 4; + const size_t vl4 = __riscv_vsetvl_e8m1(4); + for (; ii + (packn - 1) < max_ii; ii += packn) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + const int mr = packn; + const size_t vl = __riscv_vsetvl_e32m1(mr); + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + vfloat32m1_t _fsum0; + vfloat32m1_t _fsum1; + vfloat32m1_t _fsum2; + vfloat32m1_t _fsum3; + if (k == 0) + { + _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + _fsum2 = __riscv_vfmv_v_f_f32m1(0.f, vl); + _fsum3 = __riscv_vfmv_v_f_f32m1(0.f, vl); + } + else + { + _fsum0 = __riscv_vle32_v_f32m1(outptr, vl); + _fsum1 = __riscv_vle32_v_f32m1(outptr + mr, vl); + _fsum2 = __riscv_vle32_v_f32m1(outptr + mr * 2, vl); + _fsum3 = __riscv_vle32_v_f32m1(outptr + mr * 3, vl); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, 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_kk0; kk += 4) + { + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + uint32_t b = *(const uint32_t*)pB; + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint32_t*)(pB + 4); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr * 2, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint32_t*)(pB + 8); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr * 3, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint32_t*)(pB + 12); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + pA += mr * 4; + pB += 16; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + uint32_t b = *(const uint32_t*)pB; + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint32_t*)(pB + 4); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + pA += mr * 2; + pB += 8; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + const uint32_t b = *(const uint32_t*)pB; + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + pA += mr; + pB += 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_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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + vfloat32m1_t _fsum0; + vfloat32m1_t _fsum1; + if (k == 0) + { + _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + } + else + { + _fsum0 = __riscv_vle32_v_f32m1(outptr, vl); + _fsum1 = __riscv_vle32_v_f32m1(outptr + mr, vl); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, 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_kk0; kk += 4) + { + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + uint16_t b = *(const uint16_t*)pB; + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint16_t*)(pB + 2); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr * 2, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint16_t*)(pB + 4); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr * 3, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint16_t*)(pB + 6); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + pA += mr * 4; + pB += 8; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + uint16_t b = *(const uint16_t*)pB; + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + b = *(const uint16_t*)(pB + 2); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + pA += mr * 2; + pB += 4; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + const uint16_t b = *(const uint16_t*)pB; + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + pA += mr; + pB += 2; + } + + 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; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + vfloat32m1_t _fsum; + if (k == 0) + _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + else + _fsum = __riscv_vle32_v_f32m1(outptr, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + vint8m1_t _b = __riscv_vle8_v_i8m1(pB, vl4); + const signed char b0 = __riscv_vmv_x_s_i8m1_i8(_b); + _b = __riscv_vslidedown_vx_i8m1(_b, 1, vl4); + const signed char b1 = __riscv_vmv_x_s_i8m1_i8(_b); + _b = __riscv_vslidedown_vx_i8m1(_b, 1, vl4); + const signed char b2 = __riscv_vmv_x_s_i8m1_i8(_b); + _b = __riscv_vslidedown_vx_i8m1(_b, 1, vl4); + const signed char b3 = __riscv_vmv_x_s_i8m1_i8(_b); + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, b0, _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, b1, _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr * 2, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, b2, _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr * 3, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, b3, _a, vl); + pA += mr * 4; + pB += 4; + } + for (; kk + 1 < max_kk0; kk += 2) + { + const uint16_t b = *(const uint16_t*)pB; + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, (signed char)b, _a, vl); + _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA + mr, vl), 0, vl); + _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, (signed char)(b >> 8), _a, vl); + pA += mr * 2; + pB += 2; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _a16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pA, vl), 0, vl); + vint32m4_t _a32 = __riscv_vwadd_vx_i32m4(_a16, 0, vl); + vint32m1_t _a = __riscv_vget_v_i32m4_i32m1(_a32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pB[0], _a, vl); + pA += mr; + pB++; + } + + 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++; + } + + __riscv_vse32_v_f32m1(outptr, _fsum, vl); + outptr += mr; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep * mr; + pAT_descales += A_descales_hstep * mr; + } + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(4); + vfloat32m1_t _fsum0; + vfloat32m1_t _fsum1; + if (k == 0) + { + _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + } + else + { + vfloat32m1x2_t _s = __riscv_vlseg2e32_v_f32m1x2(outptr, vl); + _fsum0 = __riscv_vget_v_f32m1x2_f32m1(_s, 0); + _fsum1 = __riscv_vget_v_f32m1x2_f32m1(_s, 1); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, 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_kk0; kk += 4) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 4, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 8, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[4], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[5], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 12, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[6], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[7], _b0, vl); + pA += 8; + pB += 16; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 4, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + pA += 4; + pB += 8; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b, vl); + pA += 2; + pB += 4; + } + + 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; + } + + __riscv_vsseg2e32_v_f32m1x2(outptr, __riscv_vcreate_v_f32m1x2(_fsum0, _fsum1), vl); + outptr += 8; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(2); + vfloat32m1_t _fsum0; + vfloat32m1_t _fsum1; + if (k == 0) + { + _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + } + else + { + vfloat32m1x2_t _s = __riscv_vlseg2e32_v_f32m1x2(outptr, vl); + _fsum0 = __riscv_vget_v_f32m1x2_f32m1(_s, 0); + _fsum1 = __riscv_vget_v_f32m1x2_f32m1(_s, 1); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, 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_kk0; kk += 4) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 2, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 4, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[4], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[5], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 6, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[6], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[7], _b0, vl); + pA += 8; + pB += 8; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 2, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + pA += 4; + pB += 4; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b, vl); + pA += 2; + 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(_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; + } + + __riscv_vsseg2e32_v_f32m1x2(outptr, __riscv_vcreate_v_f32m1x2(_fsum0, _fsum1), vl); + outptr += 4; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const size_t vl = __riscv_vsetvl_e32m1(1); + vfloat32m1_t _fsum0; + vfloat32m1_t _fsum1; + if (k == 0) + { + _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + } + else + { + vfloat32m1x2_t _s = __riscv_vlseg2e32_v_f32m1x2(outptr, vl); + _fsum0 = __riscv_vget_v_f32m1x2_f32m1(_s, 0); + _fsum1 = __riscv_vget_v_f32m1x2_f32m1(_s, 1); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, 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_kk0; kk += 4) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 1, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 2, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[4], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[5], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 3, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[6], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[7], _b0, vl); + pA += 8; + pB += 4; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 1, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + pA += 4; + pB += 2; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum0 = __riscv_vmacc_vx_i32m1(_sum0, pA[0], _b, vl); + _sum1 = __riscv_vmacc_vx_i32m1(_sum1, pA[1], _b, vl); + pA += 2; + pB++; + } + + 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++; + } + + __riscv_vsseg2e32_v_f32m1x2(outptr, __riscv_vcreate_v_f32m1x2(_fsum0, _fsum1), vl); + outptr += 2; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep * 2; + pAT_descales += A_descales_hstep * 2; + } + for (; ii < max_ii; ii++) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(4); + vfloat32m1_t _fsum; + if (k == 0) + _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + else + _fsum = __riscv_vle32_v_f32m1(outptr, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 4, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 8, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[2], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 12, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[3], _b0, vl); + pA += 4; + pB += 16; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 4, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[1], _b0, vl); + pA += 2; + pB += 8; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[0], _b, vl); + pA++; + pB += 4; + } + + 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; + } + + __riscv_vse32_v_f32m1(outptr, _fsum, vl); + outptr += 4; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(2); + vfloat32m1_t _fsum; + if (k == 0) + _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + else + _fsum = __riscv_vle32_v_f32m1(outptr, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 2, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 4, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[2], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 6, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[3], _b0, vl); + pA += 4; + pB += 8; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 2, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[1], _b0, vl); + pA += 2; + pB += 4; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_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; + } + + __riscv_vse32_v_f32m1(outptr, _fsum, vl); + outptr += 2; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + const size_t vl = __riscv_vsetvl_e32m1(1); + vfloat32m1_t _fsum; + if (k == 0) + _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + else + _fsum = __riscv_vle32_v_f32m1(outptr, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 1, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[1], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 2, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[2], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 3, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[3], _b0, vl); + pA += 4; + pB += 4; + } + for (; kk + 1 < max_kk0; kk += 2) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB + 1, vl), 0, vl); + _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + _b0 = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_vx_i32m1(_sum, pA[1], _b0, vl); + pA += 2; + pB += 2; + } + for (; kk < max_kk0; kk++) + { + vint16m2_t _b16 = __riscv_vwadd_vx_i16m2(__riscv_vle8_v_i8m1(pB, vl), 0, vl); + vint32m4_t _b32 = __riscv_vwadd_vx_i32m4(_b16, 0, vl); + vint32m1_t _b = __riscv_vget_v_i32m4_i32m1(_b32, 0); + _sum = __riscv_vmacc_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++; + } + + __riscv_vse32_v_f32m1(outptr, _fsum, vl); + outptr++; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep; + pAT_descales += A_descales_hstep; + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + float sum00; + float sum01; + float sum02; + float sum03; + float sum10; + float sum11; + float sum12; + float sum13; + if (k == 0) + { + sum00 = 0.f; + sum01 = 0.f; + sum02 = 0.f; + sum03 = 0.f; + sum10 = 0.f; + sum11 = 0.f; + sum12 = 0.f; + sum13 = 0.f; + } + else + { + sum00 = outptr[0]; + sum10 = outptr[1]; + sum01 = outptr[2]; + sum11 = outptr[3]; + sum02 = outptr[4]; + sum12 = outptr[5]; + sum03 = outptr[6]; + sum13 = outptr[7]; + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, 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_kk0; 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_kk0; 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_kk0; 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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float sum00; + float sum01; + float sum10; + float sum11; + if (k == 0) + { + sum00 = 0.f; + sum01 = 0.f; + sum10 = 0.f; + sum11 = 0.f; + } + else + { + sum00 = outptr[0]; + sum10 = outptr[1]; + sum01 = outptr[2]; + sum11 = outptr[3]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + int s00 = 0, s01 = 0, s10 = 0, s11 = 0; + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0; 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_kk0; 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; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float sum0; + float sum1; + if (k == 0) + { + sum0 = 0.f; + sum1 = 0.f; + } + else + { + sum0 = outptr[0]; + sum1 = outptr[1]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + int s0 = 0, s1 = 0; + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0; 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_kk0; 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; + pB_panel += K; + pB_descales_panel += block_count; + } + pAT += A_hstep * 2; + pAT_descales += A_descales_hstep * 2; + } + for (; ii < max_ii; ii++) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)4 * k; + const float* pB_descales = pB_descales_panel + (size_t)4 * block_start; + float sum0; + float sum1; + float sum2; + float sum3; + if (k == 0) + { + sum0 = 0.f; + sum1 = 0.f; + sum2 = 0.f; + sum3 = 0.f; + } + else + { + sum0 = outptr[0]; + sum1 = outptr[1]; + sum2 = outptr[2]; + sum3 = outptr[3]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + int s0 = 0, s1 = 0, s2 = 0, s3 = 0; + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0; 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_kk0; 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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)2 * k; + const float* pB_descales = pB_descales_panel + (size_t)2 * block_start; + float sum0; + float sum1; + if (k == 0) + { + sum0 = 0.f; + sum1 = 0.f; + } + else + { + sum0 = outptr[0]; + sum1 = outptr[1]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + int s0 = 0, s1 = 0; + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0; 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_kk0; 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; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + k; + const float* pB_descales = pB_descales_panel + block_start; + float sum; + if (k == 0) + sum = 0.f; + else + sum = outptr[0]; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < A_hstep; kk0 += block_size) + { + const int max_kk0 = std::min(A_hstep - kk0, block_size); + int s = 0; + int kk = 0; + for (; kk + 3 < max_kk0; 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_kk0; kk += 2) + { + const signed char* b = pB; + s += pA[0] * b[0] + pA[1] * b[1]; + pA += 2; + pB += 2; + } + for (; kk < max_kk0; kk++) + { + s += pA[0] * pB[0]; + pA++; + pB++; + } + sum += s * pA_descales[0] * pB_descales[0]; + pA_descales++; + pB_descales++; + } + outptr[0] = sum; + outptr++; + pB_panel += K; + pB_descales_panel += block_count; + } + pAT += A_hstep; + pAT_descales += A_descales_hstep; + } +#endif // __riscv_vector +} + +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, float alpha, float beta) +{ + const float* pp = topT; + 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); + + float* out0 = outptr; + for (int jj = 0; jj < max_jj; jj++) + { + vfloat32m1_t _sum = __riscv_vle32_v_f32m1(pp, vl_packn); + pp += 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) + { + 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 * beta, vl_packn); + pC++; + } + } + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m1(_sum, alpha, vl_packn); + __riscv_vsse32_v_f32m1(out0, out_stride, _sum, vl_packn); + out0++; + } + 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) + { + 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; +#if __riscv_vector + while (jj < max_jj) + { + const size_t vl = __riscv_vsetvl_e32m4(max_jj - jj); + 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); + pp += vl * 2; + + 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) + { + vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); + vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); + pC += 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) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + 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); + } + } + } + + if (alpha != 1.f) + { + _sum0 = __riscv_vfmul_vf_f32m4(_sum0, alpha, vl); + _sum1 = __riscv_vfmul_vf_f32m4(_sum1, alpha, vl); + } + + __riscv_vse32_v_f32m4(out0, _sum0, vl); + __riscv_vse32_v_f32m4(out1, _sum1, vl); + jj += (int)vl; + out0 += vl; + out1 += vl; + } +#endif // __riscv_vector + for (; jj + 3 < max_jj; jj += 4) + { + float sum00 = pp[0]; + float sum10 = pp[1]; + float sum01 = pp[2]; + float sum11 = pp[3]; + float sum02 = pp[4]; + float sum12 = pp[5]; + float sum03 = pp[6]; + float sum13 = pp[7]; + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0 || 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; + pC += 4; + } + 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; + pC += 4; + } + } + + if (alpha != 1.f) + { + sum00 *= alpha; + sum10 *= alpha; + sum01 *= alpha; + sum11 *= alpha; + sum02 *= alpha; + sum12 *= alpha; + sum03 *= alpha; + sum13 *= alpha; + } + + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum00 = pp[0]; + float sum10 = pp[1]; + float sum01 = pp[2]; + float sum11 = pp[3]; + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0 || 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; + pC += 2; + } + if (broadcast_type_C == 4) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum10 += pC[0] * beta; + sum11 += pC[1] * beta; + pC += 2; + } + } + + if (alpha != 1.f) + { + sum00 *= alpha; + sum10 *= alpha; + sum01 *= alpha; + sum11 *= alpha; + } + + out0[0] = sum00; + out0[1] = sum01; + out1[0] = sum10; + out1[1] = sum11; + out0 += 2; + out1 += 2; + } + for (; jj < max_jj; jj++) + { + 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 += c1; + } + if (broadcast_type_C == 3) + { + sum0 += pC[0] * beta; + sum1 += pC[c_hstep] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + out0[0] = sum0; + out1[0] = sum1; + out0++; + out1++; + } + 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; +#if __riscv_vector + while (jj < max_jj) + { + const size_t vl = __riscv_vsetvl_e32m4(max_jj - jj); + vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); + pp += 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 || broadcast_type_C == 4) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + 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 (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, vl); + + __riscv_vse32_v_f32m4(out0, _sum, vl); + jj += (int)vl; + out0 += vl; + } +#endif // __riscv_vector + for (; jj + 3 < max_jj; jj += 4) + { + float sum0 = pp[0]; + float sum1 = pp[1]; + float sum2 = pp[2]; + float sum3 = pp[3]; + 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; + pC += 4; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + sum2 += pC[2] * beta; + sum3 += pC[3] * beta; + pC += 4; + } + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + sum2 *= alpha; + sum3 *= alpha; + } + out0[0] = sum0; + out0[1] = sum1; + out0[2] = sum2; + out0[3] = sum3; + out0 += 4; + } + 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) + { + 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; + } + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + out0[0] = sum0; + out0[1] = sum1; + out0 += 2; + } + for (; jj < max_jj; jj++) + { + float sum = *pp++; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + sum += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + sum += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) + sum *= alpha; + out0[0] = sum; + out0++; + } + outptr += out_hstep; + } +} + +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, float alpha, float beta) +{ + const float* pp = topT; + 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; + + 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); + + float* out0 = outptr; + for (int jj = 0; jj < max_jj; jj++) + { + vfloat32m1_t _sum = __riscv_vle32_v_f32m1(pp, vl_packn); + pp += 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) + { + 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 * beta, vl_packn); + pC++; + } + } + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m1(_sum, alpha, vl_packn); + __riscv_vse32_v_f32m1(out0, _sum, vl_packn); + out0 += out_hstep; + } + outptr += packn; + } +#endif // __riscv_vector + 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; +#if __riscv_vector + while (jj < max_jj) + { + const size_t vl = __riscv_vsetvl_e32m4(max_jj - jj); + 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); + pp += vl * 2; + 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) + { + vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); + vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); + pC += 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) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + 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); + } + } + } + if (alpha != 1.f) + { + _sum0 = __riscv_vfmul_vf_f32m4(_sum0, alpha, vl); + _sum1 = __riscv_vfmul_vf_f32m4(_sum1, alpha, vl); + } + + 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; + out0 += out_hstep * vl; + } +#endif // __riscv_vector + for (; jj + 3 < max_jj; jj += 4) + { + float sum00 = pp[0]; + float sum10 = pp[1]; + float sum01 = pp[2]; + float sum11 = pp[3]; + float sum02 = pp[4]; + float sum12 = pp[5]; + float sum03 = pp[6]; + float sum13 = pp[7]; + pp += 8; + if (pC) + { + if (broadcast_type_C == 0 || 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; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 4; + } + if (alpha != 1.f) + { + sum00 *= alpha; + sum10 *= alpha; + sum01 *= alpha; + sum11 *= alpha; + sum02 *= alpha; + sum12 *= alpha; + sum03 *= alpha; + sum13 *= alpha; + } + + 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; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum00 = pp[0]; + float sum10 = pp[1]; + float sum01 = pp[2]; + float sum11 = pp[3]; + pp += 4; + if (pC) + { + if (broadcast_type_C == 0 || 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; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 2; + } + if (alpha != 1.f) + { + sum00 *= alpha; + sum10 *= alpha; + sum01 *= alpha; + sum11 *= alpha; + } + + out0[0] = sum00; + out0[1] = sum10; + out0[out_hstep] = sum01; + out0[out_hstep + 1] = sum11; + out0 += out_hstep * 2; + } + for (; jj < max_jj; jj++) + { + 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 += 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; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC++; + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + out0[0] = sum0; + out0[1] = sum1; + out0 += out_hstep; + } + 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; +#if __riscv_vector + while (jj < max_jj) + { + const size_t vl = __riscv_vsetvl_e32m4(max_jj - jj); + vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); + pp += 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 || broadcast_type_C == 4) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + 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 (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, 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; + out0 += out_hstep * vl; + } +#endif // __riscv_vector + for (; jj + 3 < max_jj; jj += 4) + { + float sum0 = pp[0]; + float sum1 = pp[1]; + float sum2 = pp[2]; + float sum3 = pp[3]; + 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; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 4; + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + sum2 *= alpha; + sum3 *= alpha; + } + out0[0] = sum0; + out0[out_hstep] = sum1; + out0[out_hstep * 2] = sum2; + out0[out_hstep * 3] = sum3; + out0 += out_hstep * 4; + } + 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) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 2; + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + out0[0] = sum0; + out0[out_hstep] = sum1; + out0 += out_hstep * 2; + } + for (; jj < max_jj; jj++) + { + float sum = *pp++; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + sum += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + sum += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) + sum *= alpha; + out0[0] = sum; + out0 += out_hstep; + } + outptr++; + } +} + +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(); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + +#if __riscv_vector + const int tile_m_align = csrr_vlenb() / 4; + const int tile_n_align = csrr_vlenb(); +#else + const int tile_m_align = 2; + const int tile_n_align = 4; +#endif // __riscv_vector + + int tile_size = (int)sqrtf((float)l2_cache_size / (2 * sizeof(signed char) + sizeof(float))); + TILE_M = std::max(tile_m_align, tile_size / tile_m_align * tile_m_align); + TILE_N = std::max(tile_n_align, tile_size / tile_n_align * tile_n_align); + TILE_K = std::max(block_size, tile_size / block_size * block_size); + + if (K > 0) + { + const int nn_K = (K + TILE_K - 1) / TILE_K; + TILE_K = std::min(TILE_K, ((K + nn_K - 1) / nn_K + block_size - 1) / block_size * block_size); + TILE_K = std::min(TILE_K, K); + + if (nn_K == 1) + { + tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / TILE_K)); + TILE_M = std::max(tile_m_align, tile_size / tile_m_align * tile_m_align); + TILE_N = std::max(tile_n_align, tile_size / tile_n_align * tile_n_align); + } + } + + TILE_M *= std::min(nT, get_physical_cpu_count()); + + if (M > 0) + { + const int nn_M = (M + TILE_M - 1) / TILE_M; + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + tile_m_align - 1) / tile_m_align * tile_m_align); + } + + 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); + } + + if (nT > 1) + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + tile_m_align - 1) / tile_m_align * tile_m_align); + + // always take constant TILE_M/N/K value when provided + if (constant_TILE_M > 0) + TILE_M = (constant_TILE_M + tile_m_align - 1) / tile_m_align * tile_m_align; + + 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); + } +} diff --git a/src/layer/riscv/multiheadattention_riscv.cpp b/src/layer/riscv/multiheadattention_riscv.cpp new file mode 100644 index 00000000000..c6506cb297a --- /dev/null +++ b/src/layer/riscv/multiheadattention_riscv.cpp @@ -0,0 +1,709 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#include "multiheadattention_riscv.h" + +#include "layer_type.h" + +namespace ncnn { + +MultiHeadAttention_riscv::MultiHeadAttention_riscv() +{ + 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) + { + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + + if (weight_bits == 8) + return create_pipeline_wq_int8(_opt); + } +#endif + + return MultiHeadAttention::create_pipeline(_opt); +} + +int MultiHeadAttention_riscv::destroy_pipeline(const Option& _opt) +{ +#if NCNN_WEIGHT_QUANT + if (!weight_block_quantize) + return MultiHeadAttention::destroy_pipeline(_opt); + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + + if (weight_bits != 8) + return MultiHeadAttention::destroy_pipeline(_opt); +#else + return MultiHeadAttention::destroy_pipeline(_opt); +#endif + + 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; + + 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 MultiHeadAttention::destroy_pipeline(_opt); +} + +int MultiHeadAttention_riscv::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& _opt) const +{ +#if NCNN_WEIGHT_QUANT + if (!weight_block_quantize) + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + + if (weight_bits != 8) + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); +#else + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); +#endif + + 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; + 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; + + 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 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.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]; + } + + 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 00000000000..728331d735d --- /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 00000000000..86efaaf2eab --- /dev/null +++ b/src/layer/x86/gemm_wq_int8.h @@ -0,0 +1,8941 @@ +// 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, Mat& BT_tile, Mat& BT_descales_tile, 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 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, 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, 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, Mat& BT_tile, Mat& BT_descales_tile, 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 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, 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, 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, Mat& BT_tile, Mat& BT_descales_tile, 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 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, 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, 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, Mat& BT_tile, Mat& BT_descales_tile, 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 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, 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, 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 max_kk, int K, int block_size); +#endif + +static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT_tile, Mat& BT_descales_tile, 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, BT_tile, BT_descales_tile, 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, BT_tile, BT_descales_tile, 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, BT_tile, BT_descales_tile, 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, BT_tile, BT_descales_tile, j, max_jj, K, block_size); + return; + } +#endif + + const int block_count = (K + block_size - 1) / block_size; + unsigned char* pp = BT_tile; + float* pd = BT_descales_tile; + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#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 max_kk = std::min(K - g * block_size, block_size); + + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ +#if __AVXVNNIINT8__ + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _p = _mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char)); + _mm256_storeu_si256((__m256i*)pp, _p); + pp += 32; + p0 += 4; + } +#else // __AVXVNNIINT8__ + __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) + { + __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++) + { + __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++; + } + + __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 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) + { + __m128i _p = _mm_setr_epi32(((const int*)p0)[0], ((const int*)p1)[0], ((const int*)p2)[0], ((const int*)p3)[0]); +#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) + { + __m128i _p01 = _mm_setr_epi16(((const short*)p0)[0], ((const short*)p1)[0], ((const short*)p2)[0], ((const short*)p3)[0], 0, 0, 0, 0); + __m128i _p23 = _mm_setr_epi16(((const short*)(p0 + 2))[0], ((const short*)(p1 + 2))[0], ((const short*)(p2 + 2))[0], ((const short*)(p3 + 2))[0], 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__ + for (; kk + 1 < max_kk; kk += 2) + { + __m128i _p = _mm_setr_epi16(((const short*)p0)[0], ((const short*)p1)[0], ((const short*)p2)[0], ((const short*)p3)[0], 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)*p0++; + pp[1] = (unsigned char)*p1++; + pp[2] = (unsigned char)*p2++; + pp[3] = (unsigned char)*p3++; + pp += 4; + } + + pd[0] = 1.f / *ps0++; + pd[1] = 1.f / *ps1++; + pd[2] = 1.f / *ps2++; + pd[3] = 1.f / *ps3++; + pd += 4; + } + } +#endif // defined(__x86_64__) || defined(_M_X64) +#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 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) + { + __m128i _p = _mm_setr_epi32(((const int*)p0)[0], ((const int*)p1)[0], 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) + { + ((short*)pp)[0] = ((const short*)p0)[0]; + ((short*)(pp + 2))[0] = ((const short*)p1)[0]; + ((short*)(pp + 4))[0] = ((const short*)(p0 + 2))[0]; + ((short*)(pp + 6))[0] = ((const short*)(p1 + 2))[0]; + pp += 8; + p0 += 4; + p1 += 4; + } +#endif // __SSE2__ +#endif // __AVX512VNNI__ || __AVXVNNI__ +#if __SSE2__ + for (; kk + 1 < max_kk; kk += 2) + { + ((short*)pp)[0] = ((const short*)p0)[0]; + ((short*)(pp + 2))[0] = ((const short*)p1)[0]; + pp += 4; + p0 += 2; + p1 += 2; + } +#endif // __SSE2__ + for (; kk < max_kk; kk++) + { + pp[0] = (unsigned char)*p0++; + pp[1] = (unsigned char)*p1++; + pp += 2; + } + + 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 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) + { +#if !__AVXVNNIINT8__ + __m128i _p = _mm_castps_si128(_mm_load1_ps((const float*)p0)); + _p = _mm_add_epi8(_p, _mm_set1_epi8(127)); + ((int*)pp)[0] = _mm_cvtsi128_si32(_p); +#else // __AVXVNNIINT8__ + ((int*)pp)[0] = ((const int*)p0)[0]; +#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) + { +#if __SSE2__ + ((int*)pp)[0] = ((const int*)p0)[0]; +#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__ + for (; kk + 1 < max_kk; kk += 2) + { +#if __SSE2__ + ((short*)pp)[0] = ((const short*)p0)[0]; +#else + pp[0] = p0[0]; + pp[1] = p0[1]; +#endif // __SSE2__ + pp += 2; + p0 += 2; + } + for (; kk < max_kk; kk++) + { + *pp++ = (unsigned char)*p0++; + } + + pd[0] = 1.f / *ps0++; + 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) + { + 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); + + __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 _s; + if (input_scale_ptr) + _s = _mm512_loadu_ps(input_scale_ptr + k0 + kk_absmax); + + __m512 _p = _mm512_loadu_ps(p0 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax0 = _mm512_max_ps(_absmax0, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax1 = _mm512_max_ps(_absmax1, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 2 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax2 = _mm512_max_ps(_absmax2, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 3 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax3 = _mm512_max_ps(_absmax3, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 4 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax4 = _mm512_max_ps(_absmax4, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 5 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax5 = _mm512_max_ps(_absmax5, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 6 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax6 = _mm512_max_ps(_absmax6, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 7 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax7 = _mm512_max_ps(_absmax7, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 8 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax8 = _mm512_max_ps(_absmax8, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 9 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmax9 = _mm512_max_ps(_absmax9, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 10 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmaxa = _mm512_max_ps(_absmaxa, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 11 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmaxb = _mm512_max_ps(_absmaxb, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 12 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmaxc = _mm512_max_ps(_absmaxc, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 13 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmaxd = _mm512_max_ps(_absmaxd, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 14 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmaxe = _mm512_max_ps(_absmaxe, abs512_ps(_p)); + _p = _mm512_loadu_ps(p0 + A_hstep * 15 + kk_absmax); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _s); + _absmaxf = _mm512_max_ps(_absmaxf, abs512_ps(_p)); + } + + 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) + { + __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)); + } + __m512 _absmax = _mm512_setr_ps(absmax0, absmax1, absmax2, absmax3, absmax4, absmax5, absmax6, absmax7, absmax8, absmax9, absmaxa, absmaxb, absmaxc, absmaxd, absmaxe, absmaxf); + + __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(); + __m512i _v127 = _mm512_set1_epi8(127); +#endif + signed char* pp = pp0; + int kk = 0; + for (; kk + 15 < max_kk; kk += 16) + { + { + __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); + __m256 _p8 = _mm256_loadu_ps(p0 + A_hstep * 8 + kk); + __m256 _p9 = _mm256_loadu_ps(p0 + A_hstep * 9 + kk); + __m256 _pa = _mm256_loadu_ps(p0 + A_hstep * 10 + kk); + __m256 _pb = _mm256_loadu_ps(p0 + A_hstep * 11 + kk); + __m256 _pc = _mm256_loadu_ps(p0 + A_hstep * 12 + kk); + __m256 _pd = _mm256_loadu_ps(p0 + A_hstep * 13 + kk); + __m256 _pe = _mm256_loadu_ps(p0 + A_hstep * 14 + kk); + __m256 _pf = _mm256_loadu_ps(p0 + A_hstep * 15 + kk); + if (input_scale_ptr) + { + __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); + _p8 = _mm256_mul_ps(_p8, _s); + _p9 = _mm256_mul_ps(_p9, _s); + _pa = _mm256_mul_ps(_pa, _s); + _pb = _mm256_mul_ps(_pb, _s); + _pc = _mm256_mul_ps(_pc, _s); + _pd = _mm256_mul_ps(_pd, _s); + _pe = _mm256_mul_ps(_pe, _s); + _pf = _mm256_mul_ps(_pf, _s); + } + transpose8x8_ps(_p0, _p1, _p2, _p3, _p4, _p5, _p6, _p7); + transpose8x8_ps(_p8, _p9, _pa, _pb, _pc, _pd, _pe, _pf); + + __m512 _t0 = combine8x2_ps(_p0, _p8); + __m512 _t1 = combine8x2_ps(_p1, _p9); + __m512 _t2 = combine8x2_ps(_p2, _pa); + __m512 _t3 = combine8x2_ps(_p3, _pb); + __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); + __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _v127, _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; + + _t0 = combine8x2_ps(_p4, _pc); + _t1 = combine8x2_ps(_p5, _pd); + _t2 = combine8x2_ps(_p6, _pe); + _t3 = combine8x2_ps(_p7, _pf); + _q0 = float2int8_avx512(_mm512_mul_ps(_t0, _scale)); + _q1 = float2int8_avx512(_mm512_mul_ps(_t1, _scale)); + _q2 = float2int8_avx512(_mm512_mul_ps(_t2, _scale)); + _q3 = float2int8_avx512(_mm512_mul_ps(_t3, _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, _v127, _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; + } + { + __m256 _p0 = _mm256_loadu_ps(p0 + kk + 8); + __m256 _p1 = _mm256_loadu_ps(p0 + A_hstep + kk + 8); + __m256 _p2 = _mm256_loadu_ps(p0 + A_hstep * 2 + kk + 8); + __m256 _p3 = _mm256_loadu_ps(p0 + A_hstep * 3 + kk + 8); + __m256 _p4 = _mm256_loadu_ps(p0 + A_hstep * 4 + kk + 8); + __m256 _p5 = _mm256_loadu_ps(p0 + A_hstep * 5 + kk + 8); + __m256 _p6 = _mm256_loadu_ps(p0 + A_hstep * 6 + kk + 8); + __m256 _p7 = _mm256_loadu_ps(p0 + A_hstep * 7 + kk + 8); + __m256 _p8 = _mm256_loadu_ps(p0 + A_hstep * 8 + kk + 8); + __m256 _p9 = _mm256_loadu_ps(p0 + A_hstep * 9 + kk + 8); + __m256 _pa = _mm256_loadu_ps(p0 + A_hstep * 10 + kk + 8); + __m256 _pb = _mm256_loadu_ps(p0 + A_hstep * 11 + kk + 8); + __m256 _pc = _mm256_loadu_ps(p0 + A_hstep * 12 + kk + 8); + __m256 _pd = _mm256_loadu_ps(p0 + A_hstep * 13 + kk + 8); + __m256 _pe = _mm256_loadu_ps(p0 + A_hstep * 14 + kk + 8); + __m256 _pf = _mm256_loadu_ps(p0 + A_hstep * 15 + kk + 8); + if (input_scale_ptr) + { + __m256 _s = _mm256_loadu_ps(input_scale_ptr + k0 + kk + 8); + _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); + _p8 = _mm256_mul_ps(_p8, _s); + _p9 = _mm256_mul_ps(_p9, _s); + _pa = _mm256_mul_ps(_pa, _s); + _pb = _mm256_mul_ps(_pb, _s); + _pc = _mm256_mul_ps(_pc, _s); + _pd = _mm256_mul_ps(_pd, _s); + _pe = _mm256_mul_ps(_pe, _s); + _pf = _mm256_mul_ps(_pf, _s); + } + transpose8x8_ps(_p0, _p1, _p2, _p3, _p4, _p5, _p6, _p7); + transpose8x8_ps(_p8, _p9, _pa, _pb, _pc, _pd, _pe, _pf); + + __m512 _t0 = combine8x2_ps(_p0, _p8); + __m512 _t1 = combine8x2_ps(_p1, _p9); + __m512 _t2 = combine8x2_ps(_p2, _pa); + __m512 _t3 = combine8x2_ps(_p3, _pb); + __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); + __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _v127, _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; + + _t0 = combine8x2_ps(_p4, _pc); + _t1 = combine8x2_ps(_p5, _pd); + _t2 = combine8x2_ps(_p6, _pe); + _t3 = combine8x2_ps(_p7, _pf); + _q0 = float2int8_avx512(_mm512_mul_ps(_t0, _scale)); + _q1 = float2int8_avx512(_mm512_mul_ps(_t1, _scale)); + _q2 = float2int8_avx512(_mm512_mul_ps(_t2, _scale)); + _q3 = float2int8_avx512(_mm512_mul_ps(_t3, _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, _v127, _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) + { + __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); + } + __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); + __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _v127, _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, 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])); + } + __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, p0 + kk, sizeof(float)); + if (input_scale_ptr) + { + _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); + } + _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) + { + 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); + + __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 _s; + if (input_scale_ptr) + _s = _mm256_loadu_ps(input_scale_ptr + k0 + kk); + + __m256 _p = _mm256_loadu_ps(p0 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax0 = _mm256_max_ps(_absmax0, abs256_ps(_p)); + _p = _mm256_loadu_ps(p0 + A_hstep + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax1 = _mm256_max_ps(_absmax1, abs256_ps(_p)); + _p = _mm256_loadu_ps(p0 + A_hstep * 2 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax2 = _mm256_max_ps(_absmax2, abs256_ps(_p)); + _p = _mm256_loadu_ps(p0 + A_hstep * 3 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax3 = _mm256_max_ps(_absmax3, abs256_ps(_p)); + _p = _mm256_loadu_ps(p0 + A_hstep * 4 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax4 = _mm256_max_ps(_absmax4, abs256_ps(_p)); + _p = _mm256_loadu_ps(p0 + A_hstep * 5 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax5 = _mm256_max_ps(_absmax5, abs256_ps(_p)); + _p = _mm256_loadu_ps(p0 + A_hstep * 6 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax6 = _mm256_max_ps(_absmax6, abs256_ps(_p)); + _p = _mm256_loadu_ps(p0 + A_hstep * 7 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _s); + _absmax7 = _mm256_max_ps(_absmax7, abs256_ps(_p)); + } + + 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)); + } + + __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(); +#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); + __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) + { + __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); + } + + __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); + __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(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])); + } + _p0 = _mm256_mul_ps(_p0, _scale); + _p1 = _mm256_mul_ps(_p1, _scale); + __m128i _q = float2int8_avx(_p0, _p1); + __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(p0 + kk, _vindex, sizeof(float)); + if (input_scale_ptr) + { + _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); + } + *(int64_t*)pp = 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) + { + 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); + + __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) + { + __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)); + } + + __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(); +#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) + { + __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); + } + __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__ + __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 + __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; + } +#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])); + } + __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_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])); + } + ((int*)pp)[0] = 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__ +#if __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { + 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); + + __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) + { + __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)); + } + + __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(); +#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) + { + __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm_mul_ps(_p0, _s); + _p1 = _mm_mul_ps(_p1, _s); + } + __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__ + __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])); + } + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + ((int*)pp)[0] = _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])); + } + ((short*)pp)[0] = (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 + } + } +#endif // __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { + 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); + float absmax0 = 0.f; + float absmax1 = 0.f; + for (int kk = 0; kk < max_kk; 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; + } + 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) + { + scale0 = 127.f / absmax0; + } + if (absmax1 != 0.f) + { + scale1 = 127.f / absmax1; + } + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + + signed char* pp = pp0; + for (int kk = 0; kk < max_kk; 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; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } + + p0 += max_kk; + pp0 += max_kk * 2; + pd += 2; + } + } + for (; ii < max_ii; ii++) + { + 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); + signed char* pp = pp0; + 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(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)); + } + 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(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)); + } + 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(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)); + } + absmax = std::max(absmax, _mm_reduce_max_ps(_absmax128)); +#endif // __SSE2__ + for (; kk < max_kk; kk++) + { + float v = p0[kk]; + if (input_scale_ptr) + v *= input_scale_ptr[k0 + kk]; + absmax = std::max(absmax, (float)fabsf(v)); + } + + if (absmax == 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; + } + + const float scale = 127.f / absmax; + pd[0] = absmax / 127.f; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + int w_shift = 0; +#endif + kk = 0; +#if __SSE2__ +#if __AVX__ +#if __AVX512F__ + __m512 _scale512 = _mm512_set1_ps(scale); + for (; kk + 15 < max_kk; kk += 16) + { + __m512 _p = _mm512_loadu_ps(p0 + kk); + if (input_scale_ptr) + { + _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); + } + __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); + _mm_storeu_si128((__m128i*)pp, _q); + pp += 16; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __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__ + __m256 _scale256 = _mm256_set1_ps(scale); + for (; kk + 7 < max_kk; kk += 8) + { + __m256 _p = _mm256_loadu_ps(p0 + kk); + if (input_scale_ptr) + { + _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); + } + 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) + __m128i _q8 = _mm_cvtsi64_si128(q); +#else + __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); +#endif + __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__ + __m128 _scale128 = _mm_set1_ps(scale); + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p = _mm_loadu_ps(p0 + kk); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); + } + const int32_t q = float2int8_sse(_mm_mul_ps(_p, _scale128)); + ((int*)pp)[0] = q; + pp += 4; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __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 + } +#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 = p0[kk]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k0 + kk]; + } + *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 + } + } +} + +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) + { + 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); + + __m512 _absmax = _mm512_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __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)); + } + + __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(); + __m512i _v127 = _mm512_set1_epi8(127); +#endif + signed char* pp = pp0; + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + __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])); + } + __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); + __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _v127, _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(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])); + } + __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(p0 + kk * A_hstep); + if (input_scale_ptr) + { + _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); + } + _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) + { + 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); + + __m256 _absmax = _mm256_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __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)); + } + + __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(); +#endif + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __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])); + } + _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); + __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(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])); + } + _p0 = _mm256_mul_ps(_p0, _scale); + _p1 = _mm256_mul_ps(_p1, _scale); + __m128i _q = float2int8_avx(_p0, _p1); + __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(p0 + kk * A_hstep); + if (input_scale_ptr) + { + _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); + } + *(int64_t*)pp = 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) + { + 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); + + __m128 _absmax = _mm_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __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)); + } + + __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(); +#endif + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __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])); + } + __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__ + __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(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])); + } + __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(p0 + kk * A_hstep); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); + } + ((int*)pp)[0] = 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__ +#if __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { + 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); + + __m128 _absmax = _mm_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __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)); + } + + __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(); +#endif + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __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])); + } + __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__ + __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*)(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])); + } + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + ((int*)pp)[0] = _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*)(p0 + kk * A_hstep)); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); + } + ((short*)pp)[0] = (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 + } + } +#endif // __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { + 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); + float absmax0 = 0.f; + float absmax1 = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const float* ptrA = p0 + kk * A_hstep; + 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) + { + scale0 = 127.f / absmax0; + } + if (absmax1 != 0.f) + { + scale1 = 127.f / absmax1; + } + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + + signed char* pp = pp0; + for (int kk = 0; kk < max_kk; kk++) + { + const float* ptrA = p0 + kk * A_hstep; + float v0 = ptrA[0]; + float v1 = ptrA[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k0 + kk]; + v0 *= s; + v1 *= s; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } + + p0 += max_kk * A_hstep; + pp0 += max_kk * 2; + pd += 2; + } + } + +#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++) + { + 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); + signed char* pp = pp0; + 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 = 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)); + _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 = 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)); + _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 = 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)); + _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 = p0[kk * A_hstep]; + if (input_scale_ptr) + v *= input_scale_ptr[k0 + kk]; + absmax = std::max(absmax, (float)fabsf(v)); + } + + if (absmax == 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; + } + + const float scale = 127.f / absmax; + pd[0] = absmax / 127.f; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + int w_shift = 0; +#endif + kk = 0; +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + __m512 _scale512 = _mm512_set1_ps(scale); + for (; kk + 15 < max_kk; kk += 16) + { + 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)); + } + __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); + _mm_storeu_si128((__m128i*)pp, _q); + pp += 16; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __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__ + __m256 _scale256 = _mm256_set1_ps(scale); + for (; kk + 7 < max_kk; kk += 8) + { + 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)); + } + 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) + __m128i _q8 = _mm_cvtsi64_si128(q); +#else + __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); +#endif + __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__ + __m128 _scale128 = _mm_set1_ps(scale); + for (; kk + 3 < max_kk; kk += 4) + { + 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)); + } + const int32_t q = float2int8_sse(_mm_mul_ps(_p, _scale128)); + ((int*)pp)[0] = q; + pp += 4; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __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 + } +#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 = p0[kk * A_hstep]; + if (input_scale_ptr) + { + 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 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, 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, 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, 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, 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, max_kk, 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; + const int tile_K = max_kk; + const int block_count = (K + block_size - 1) / block_size; + const int block_start = k / block_size; + + int ii = 0; +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + for (; ii + 15 < max_ii; ii += 16) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)k * 8; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 8; + __m512 _fsum0; + __m512 _fsum1; + __m512 _fsum2; + __m512 _fsum3; + __m512 _fsum4; + __m512 _fsum5; + __m512 _fsum6; + __m512 _fsum7; + + if (k == 0) + { + _fsum0 = _mm512_setzero_ps(); + _fsum1 = _mm512_setzero_ps(); + _fsum2 = _mm512_setzero_ps(); + _fsum3 = _mm512_setzero_ps(); + _fsum4 = _mm512_setzero_ps(); + _fsum5 = _mm512_setzero_ps(); + _fsum6 = _mm512_setzero_ps(); + _fsum7 = _mm512_setzero_ps(); + } + else + { + _fsum0 = _mm512_loadu_ps(outptr); + _fsum1 = _mm512_loadu_ps(outptr + 16); + _fsum2 = _mm512_loadu_ps(outptr + 32); + _fsum3 = _mm512_loadu_ps(outptr + 48); + _fsum4 = _mm512_loadu_ps(outptr + 64); + _fsum5 = _mm512_loadu_ps(outptr + 80); + _fsum6 = _mm512_loadu_ps(outptr + 96); + _fsum7 = _mm512_loadu_ps(outptr + 112); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + 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(); + __m512i _sum4 = _mm512_setzero_si512(); + __m512i _sum5 = _mm512_setzero_si512(); + __m512i _sum6 = _mm512_setzero_si512(); + __m512i _sum7 = _mm512_setzero_si512(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; + + // from + // 00 01 02 03 04 05 06 07 + // 10 11 12 13 14 15 16 17 + // ... + // f0 f1 f2 f3 f4 f5 f6 f7 + // + // to + // _sum0 00 11 22 33 44 55 66 77 80 91 a2 b3 c4 d5 e6 f7 + // _sum1 01 12 23 30 45 56 67 74 81 92 a3 b0 c5 d6 e7 f4 + // _sum2 20 31 02 13 64 75 46 57 a0 b1 82 93 e4 f5 c6 d7 + // _sum3 21 32 03 10 65 76 47 54 a1 b2 83 90 e5 f6 c7 d4 + // _sum4 04 15 26 37 40 51 62 73 84 95 a6 b7 c0 d1 e2 f3 + // _sum5 05 16 27 34 41 52 63 70 85 96 a7 b4 c1 d2 e3 f0 + // _sum6 24 35 06 17 60 71 42 53 a4 b5 86 97 e0 f1 c2 d3 + // _sum7 25 36 07 14 61 72 43 50 a5 b6 87 94 e1 f2 c3 d0 +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 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); + _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_kk0 >= 4) + { + __m512i _w_shift0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _w_shift1 = _mm512_alignr_epi8(_w_shift0, _w_shift0, 8); + _sum0 = _mm512_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm512_sub_epi32(_sum1, _w_shift0); + _sum2 = _mm512_sub_epi32(_sum2, _w_shift1); + _sum3 = _mm512_sub_epi32(_sum3, _w_shift1); + _sum4 = _mm512_sub_epi32(_sum4, _w_shift0); + _sum5 = _mm512_sub_epi32(_sum5, _w_shift0); + _sum6 = _mm512_sub_epi32(_sum6, _w_shift1); + _sum7 = _mm512_sub_epi32(_sum7, _w_shift1); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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); + _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_kk0; kk++) + { + __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); + __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))); + _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; + } + + __m512 _ad0 = _mm512_loadu_ps(pA_descales); + __m512 _ad1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_ad0), _mm512_castps_si512(_ad0), 8)); + __m256 _b = _mm256_loadu_ps(pB_descales); + __m512 _bd0 = combine8x2_ps(_b, _b); + __m512 _bd1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_bd0), _mm512_castps_si512(_bd0), 4)); + __m512 _bd2 = _mm512_castsi512_ps(_mm512_permutex_epi64(_mm512_castps_si512(_bd0), _MM_SHUFFLE(1, 0, 3, 2))); + __m512 _bd3 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_bd2), _mm512_castps_si512(_bd2), 4)); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_ad0, _bd0))); + _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_ad0, _bd1))); + _fsum2 = _mm512_add_ps(_fsum2, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum2), _mm512_mul_ps(_ad1, _bd0))); + _fsum3 = _mm512_add_ps(_fsum3, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum3), _mm512_mul_ps(_ad1, _bd1))); + _fsum4 = _mm512_add_ps(_fsum4, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum4), _mm512_mul_ps(_ad0, _bd2))); + _fsum5 = _mm512_add_ps(_fsum5, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum5), _mm512_mul_ps(_ad0, _bd3))); + _fsum6 = _mm512_add_ps(_fsum6, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum6), _mm512_mul_ps(_ad1, _bd2))); + _fsum7 = _mm512_add_ps(_fsum7, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum7), _mm512_mul_ps(_ad1, _bd3))); + 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; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)k * 4; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 4; + __m512 _fsum0; + __m512 _fsum1; + __m512 _fsum2; + __m512 _fsum3; + + if (k == 0) + { + _fsum0 = _mm512_setzero_ps(); + _fsum1 = _mm512_setzero_ps(); + _fsum2 = _mm512_setzero_ps(); + _fsum3 = _mm512_setzero_ps(); + } + else + { + _fsum0 = _mm512_loadu_ps(outptr); + _fsum1 = _mm512_loadu_ps(outptr + 16); + _fsum2 = _mm512_loadu_ps(outptr + 32); + _fsum3 = _mm512_loadu_ps(outptr + 48); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + 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_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 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); + _sum3 = _mm512_dpbusd_epi32(_sum3, _pB1, _pA1); + pB += 16; + pA += 64; + } + if (max_kk0 >= 4) + { + __m512i _w_shift0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _w_shift1 = _mm512_alignr_epi8(_w_shift0, _w_shift0, 8); + _sum0 = _mm512_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm512_sub_epi32(_sum1, _w_shift0); + _sum2 = _mm512_sub_epi32(_sum2, _w_shift1); + _sum3 = _mm512_sub_epi32(_sum3, _w_shift1); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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); + _sum3 = _mm512_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pB += 8; + pA += 32; + } + for (; kk < max_kk0; kk++) + { + __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))); + _sum3 = _mm512_add_epi32(_sum3, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB1))); + pB += 4; + pA += 16; + } + + __m512 _ad0 = _mm512_loadu_ps(pA_descales); + __m512 _ad1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_ad0), _mm512_castps_si512(_ad0), 8)); + __m512 _bd0 = _mm512_broadcast_f32x4(_mm_loadu_ps(pB_descales)); + __m512 _bd1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_bd0), _mm512_castps_si512(_bd0), 4)); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_ad0, _bd0))); + _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_ad0, _bd1))); + _fsum2 = _mm512_add_ps(_fsum2, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum2), _mm512_mul_ps(_ad1, _bd0))); + _fsum3 = _mm512_add_ps(_fsum3, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum3), _mm512_mul_ps(_ad1, _bd1))); + 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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)k * 2; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 2; + __m512 _fsum0; + __m512 _fsum1; + + if (k == 0) + { + _fsum0 = _mm512_setzero_ps(); + _fsum1 = _mm512_setzero_ps(); + } + else + { + _fsum0 = _mm512_loadu_ps(outptr); + _fsum1 = _mm512_loadu_ps(outptr + 16); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m512i _sum0 = _mm512_setzero_si512(); + __m512i _sum1 = _mm512_setzero_si512(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _pB0 = _mm512_set1_epi64(((const int64_t*)pB)[0]); + __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_kk0 >= 4) + { + __m512i _w_shift0 = _mm512_loadu_si512((const __m512i*)pA); + _sum0 = _mm512_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm512_sub_epi32(_sum1, _w_shift0); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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; + pA += 32; + } + for (; kk < max_kk0; kk++) + { + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m128i _pB = _mm_set1_epi16(((const short*)pB)[0]); + __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; + } + + __m512 _ad0 = _mm512_loadu_ps(pA_descales); + __m512 _bd0 = _mm512_castsi512_ps(_mm512_set1_epi64(((const int64_t*)pB_descales)[0])); + __m512 _bd1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_bd0), _mm512_castps_si512(_bd0), 4)); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_ad0, _bd0))); + _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_ad0, _bd1))); + pA_descales += 16; + pB_descales += 2; + } + + _mm512_storeu_ps(outptr + 0, _fsum0); + _mm512_storeu_ps(outptr + 16, _fsum1); + outptr += 32; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + (size_t)k; + const float* pB_descales = pB_descales_panel + (size_t)block_start; + __m512 _fsum0; + + if (k == 0) + { + _fsum0 = _mm512_setzero_ps(); + } + else + { + _fsum0 = _mm512_loadu_ps(outptr); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m512i _sum0 = _mm512_setzero_si512(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _pB0 = _mm512_set1_epi32(((const int*)pB)[0]); + _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); + pB += 4; + pA += 64; + } + if (max_kk0 >= 4) + { + __m512i _w_shift0 = _mm512_loadu_si512((const __m512i*)pA); + _sum0 = _mm512_sub_epi32(_sum0, _w_shift0); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + __m512i _pB0 = _mm512_cvtepi8_epi16(_mm256_set1_epi16(((const short*)pB)[0])); + _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + pB += 2; + pA += 32; + } + for (; kk < max_kk0; kk++) + { + __m512i _pA0 = _mm512_cvtepi8_epi32(_mm_loadu_si128((const __m128i*)pA)); + _sum0 = _mm512_add_epi32(_sum0, _mm512_mullo_epi32(_pA0, _mm512_set1_epi32(pB[0]))); + pB += 1; + pA += 16; + } + + __m512 _ad0 = _mm512_loadu_ps(pA_descales); + __m512 _bd0 = _mm512_set1_ps(pB_descales[0]); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_ad0, _bd0))); + pA_descales += 16; + pB_descales += 1; + } + + _mm512_storeu_ps(outptr + 0, _fsum0); + outptr += 16; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep * 16; + pAT_descales += A_descales_hstep * 16; + } +#endif // __AVX512F__ + for (; ii + 7 < max_ii; ii += 8) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)k * 8; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 8; + __m256 _fsum0; + __m256 _fsum1; + __m256 _fsum2; + __m256 _fsum3; + __m256 _fsum4; + __m256 _fsum5; + __m256 _fsum6; + __m256 _fsum7; + + if (k == 0) + { + _fsum0 = _mm256_setzero_ps(); + _fsum1 = _mm256_setzero_ps(); + _fsum2 = _mm256_setzero_ps(); + _fsum3 = _mm256_setzero_ps(); + _fsum4 = _mm256_setzero_ps(); + _fsum5 = _mm256_setzero_ps(); + _fsum6 = _mm256_setzero_ps(); + _fsum7 = _mm256_setzero_ps(); + } + else + { + _fsum0 = _mm256_loadu_ps(outptr); + _fsum1 = _mm256_loadu_ps(outptr + 8); + _fsum2 = _mm256_loadu_ps(outptr + 16); + _fsum3 = _mm256_loadu_ps(outptr + 24); + _fsum4 = _mm256_loadu_ps(outptr + 32); + _fsum5 = _mm256_loadu_ps(outptr + 40); + _fsum6 = _mm256_loadu_ps(outptr + 48); + _fsum7 = _mm256_loadu_ps(outptr + 56); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + 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(); + __m256i _sum4 = _mm256_setzero_si256(); + __m256i _sum5 = _mm256_setzero_si256(); + __m256i _sum6 = _mm256_setzero_si256(); + __m256i _sum7 = _mm256_setzero_si256(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 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); + _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_kk0 >= 4) + { + __m256i _w_shift0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _w_shift1 = _mm256_alignr_epi8(_w_shift0, _w_shift0, 8); + _sum0 = _mm256_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm256_sub_epi32(_sum1, _w_shift0); + _sum2 = _mm256_sub_epi32(_sum2, _w_shift1); + _sum3 = _mm256_sub_epi32(_sum3, _w_shift1); + _sum4 = _mm256_sub_epi32(_sum4, _w_shift0); + _sum5 = _mm256_sub_epi32(_sum5, _w_shift0); + _sum6 = _mm256_sub_epi32(_sum6, _w_shift1); + _sum7 = _mm256_sub_epi32(_sum7, _w_shift1); + pA += 32; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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); + _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_kk0) + { + __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); + __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))); + _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; + } + + __m256 _ad0 = _mm256_loadu_ps(pA_descales); + __m256 _ad1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_ad0), _mm256_castps_si256(_ad0), 8)); + __m256 _bd0 = _mm256_loadu_ps(pB_descales); + __m256 _bd1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_bd0), _mm256_castps_si256(_bd0), 4)); + __m256 _bd2 = _mm256_castsi256_ps(_mm256_permute4x64_epi64(_mm256_castps_si256(_bd0), _MM_SHUFFLE(1, 0, 3, 2))); + __m256 _bd3 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_bd2), _mm256_castps_si256(_bd2), 4)); + _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))); + _fsum4 = _mm256_add_ps(_fsum4, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum4), _mm256_mul_ps(_ad0, _bd2))); + _fsum5 = _mm256_add_ps(_fsum5, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum5), _mm256_mul_ps(_ad0, _bd3))); + _fsum6 = _mm256_add_ps(_fsum6, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum6), _mm256_mul_ps(_ad1, _bd2))); + _fsum7 = _mm256_add_ps(_fsum7, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum7), _mm256_mul_ps(_ad1, _bd3))); + 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; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)k * 4; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 4; + __m256 _fsum0; + __m256 _fsum1; + __m256 _fsum2; + __m256 _fsum3; + + if (k == 0) + { + _fsum0 = _mm256_setzero_ps(); + _fsum1 = _mm256_setzero_ps(); + _fsum2 = _mm256_setzero_ps(); + _fsum3 = _mm256_setzero_ps(); + } + else + { + _fsum0 = _mm256_loadu_ps(outptr); + _fsum1 = _mm256_loadu_ps(outptr + 8); + _fsum2 = _mm256_loadu_ps(outptr + 16); + _fsum3 = _mm256_loadu_ps(outptr + 24); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + 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_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 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); + _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_kk0 >= 4) + { + __m256i _w_shift0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _w_shift1 = _mm256_alignr_epi8(_w_shift0, _w_shift0, 8); + _sum0 = _mm256_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm256_sub_epi32(_sum1, _w_shift0); + _sum2 = _mm256_sub_epi32(_sum2, _w_shift1); + _sum3 = _mm256_sub_epi32(_sum3, _w_shift1); + pA += 32; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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); + _sum3 = _mm256_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pB += 8; + pA += 16; + } + for (; kk < max_kk0; 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); + __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))); + _sum3 = _mm256_add_epi32(_sum3, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB1))); + pB += 4; + pA += 8; + } + + __m256 _ad0 = _mm256_loadu_ps(pA_descales); + __m256 _ad1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_ad0), _mm256_castps_si256(_ad0), 8)); + __m128 _b = _mm_loadu_ps(pB_descales); + __m256 _bd0 = combine4x2_ps(_b, _b); + __m256 _bd1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_bd0), _mm256_castps_si256(_bd0), 4)); + _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 += 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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)k * 2; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 2; + __m256 _fsum0; + __m256 _fsum1; + + if (k == 0) + { + _fsum0 = _mm256_setzero_ps(); + _fsum1 = _mm256_setzero_ps(); + } + else + { + _fsum0 = _mm256_loadu_ps(outptr); + _fsum1 = _mm256_loadu_ps(outptr + 8); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + __m256i _sum1 = _mm256_setzero_si256(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 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); +#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_kk0 >= 4) + { + __m256i _w_shift0 = _mm256_loadu_si256((const __m256i*)pA); + _sum0 = _mm256_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm256_sub_epi32(_sum1, _w_shift0); + pA += 32; + } +#endif +#else + for (; kk + 3 < max_kk0; kk += 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); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA23, _pB23_1); + pB += 8; + pA += 32; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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; + pA += 16; + } + for (; kk < max_kk0; kk++) + { + __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB0 = _mm_set1_epi16(((const short*)pB)[0]); + _pA = _mm_cvtepi8_epi16(_pA); + _pB0 = _mm_cvtepi8_epi16(_pB0); + __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; + } + + __m256 _ad0 = _mm256_loadu_ps(pA_descales); + __m256 _bd0 = _mm256_castpd_ps(_mm256_broadcast_sd((const double*)pB_descales)); + __m256 _bd1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_bd0), _mm256_castps_si256(_bd0), 4)); + _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))); + pA_descales += 8; + pB_descales += 2; + } + + _mm256_storeu_ps(outptr + 0, _fsum0); + _mm256_storeu_ps(outptr + 8, _fsum1); + outptr += 16; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + (size_t)k; + const float* pB_descales = pB_descales_panel + (size_t)block_start; + __m256 _fsum0; + + if (k == 0) + { + _fsum0 = _mm256_setzero_ps(); + } + else + { + _fsum0 = _mm256_loadu_ps(outptr); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __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__ + _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); +#endif // __AVXVNNIINT8__ + pB += 4; + pA += 32; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk0 >= 4) + { + __m256i _w_shift0 = _mm256_loadu_si256((const __m256i*)pA); + _sum0 = _mm256_sub_epi32(_sum0, _w_shift0); + pA += 32; + } +#endif +#else + for (; kk + 3 < max_kk0; kk += 4) + { + __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_castps_si128(_mm_load1_ps((const float*)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; + pA += 32; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m256i _pB0 = _mm256_cvtepi8_epi16(_mm_set1_epi16(((const short*)pB)[0])); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + pB += 2; + pA += 16; + } + for (; kk < max_kk0; 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(pB[0])))); + pB += 1; + pA += 8; + } + + __m256 _ad0 = _mm256_loadu_ps(pA_descales); + _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_ad0, _mm256_set1_ps(pB_descales[0])))); + pA_descales += 8; + pB_descales++; + } + + _mm256_storeu_ps(outptr, _fsum0); + outptr += 8; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep * 8; + pAT_descales += A_descales_hstep * 8; + } +#endif // __AVX2__ + for (; ii + 3 < max_ii; ii += 4) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)k * 8; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 8; + __m256 _fsum0; + __m256 _fsum1; + __m256 _fsum2; + __m256 _fsum3; + + if (k == 0) + { + _fsum0 = _mm256_setzero_ps(); + _fsum1 = _mm256_setzero_ps(); + _fsum2 = _mm256_setzero_ps(); + _fsum3 = _mm256_setzero_ps(); + } + else + { + _fsum0 = _mm256_loadu_ps(outptr); + _fsum1 = _mm256_loadu_ps(outptr + 8); + _fsum2 = _mm256_loadu_ps(outptr + 16); + _fsum3 = _mm256_loadu_ps(outptr + 24); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + 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_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 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); + _sum3 = _mm256_comp_dpbusd_epi32(_sum3, _pB1, _pA1); + pA += 16; + pB += 32; + } + if (max_kk0 >= 4) + { + __m256i _w_shift0 = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)pA)); + __m256i _w_shift1 = _mm256_alignr_epi8(_w_shift0, _w_shift0, 8); + _sum0 = _mm256_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm256_sub_epi32(_sum1, _w_shift0); + _sum2 = _mm256_sub_epi32(_sum2, _w_shift1); + _sum3 = _mm256_sub_epi32(_sum3, _w_shift1); + pA += 16; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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); + _sum3 = _mm256_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pA += 8; + pB += 16; + } + for (; kk < max_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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)); + _sum3 = _mm256_add_epi32(_sum3, _mm256_mullo_epi32(_pA1, _pB1)); + pA += 4; + pB += 8; + } + + __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))); + _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; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)k * 4; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 4; + __m128 _fsum0; + __m128 _fsum1; + __m128 _fsum2; + __m128 _fsum3; + + if (k == 0) + { + _fsum0 = _mm_setzero_ps(); + _fsum1 = _mm_setzero_ps(); + _fsum2 = _mm_setzero_ps(); + _fsum3 = _mm_setzero_ps(); + } + else + { + _fsum0 = _mm_loadu_ps(outptr); + _fsum1 = _mm_loadu_ps(outptr + 4); + _fsum2 = _mm_loadu_ps(outptr + 8); + _fsum3 = _mm_loadu_ps(outptr + 12); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + 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_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 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); + _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_kk0 >= 4) + { + __m128i _w_shift0 = _mm_loadu_si128((const __m128i*)pA); + __m128i _w_shift1 = _mm_alignr_epi8(_w_shift0, _w_shift0, 8); + _sum0 = _mm_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm_sub_epi32(_sum1, _w_shift0); + _sum2 = _mm_sub_epi32(_sum2, _w_shift1); + _sum3 = _mm_sub_epi32(_sum3, _w_shift1); + pA += 16; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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); + _sum3 = _mm_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pA += 8; + pB += 8; + } + for (; kk < max_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __m128i _pB8 = _mm_cvtsi32_si128(((const int*)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)); + __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); + _sum3 = _mm_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pA += 4; + pB += 4; + } + + __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))); + _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; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)k * 2; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 2; + __m128 _fsum0; + __m128 _fsum1; + + if (k == 0) + { + _fsum0 = _mm_setzero_ps(); + _fsum1 = _mm_setzero_ps(); + } + else + { + _fsum0 = _mm_loadu_ps(outptr); + _fsum1 = _mm_loadu_ps(outptr + 4); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum0 = _mm_setzero_si128(); + __m128i _sum1 = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 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); +#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_kk0 >= 4) + { + __m128i _w_shift0 = _mm_loadu_si128((const __m128i*)pA); + _sum0 = _mm_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm_sub_epi32(_sum1, _w_shift0); + pA += 16; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __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; + pB += 4; + } + for (; kk < max_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __m128i _pB8 = _mm_set1_epi16(((const short*)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)); + __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; + } + + __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; + } + + _mm_storeu_ps(outptr, _fsum0); + _mm_storeu_ps(outptr + 4, _fsum1); + outptr += 8; + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + (size_t)k; + const float* pB_descales = pB_descales_panel + (size_t)block_start; + __m128 _fsum; + + if (k == 0) + { + _fsum = _mm_setzero_ps(); + } + else + { + _fsum = _mm_loadu_ps(outptr); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m128i _pB = _mm_set1_epi32(((const int*)pB)[0]); +#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_kk0 >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_loadu_si128((const __m128i*)pA)); + pA += 16; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB8 = _mm_set1_epi16(((const short*)pB)[0]); + __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_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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(pB[0])); + pA += 4; + pB++; + } + + __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; + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep * 4; + pAT_descales += A_descales_hstep * 4; + } +#endif // __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)k * 8; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 8; + __m256 _fsum0; + __m256 _fsum1; + + if (k == 0) + { + _fsum0 = _mm256_setzero_ps(); + _fsum1 = _mm256_setzero_ps(); + } + else + { + _fsum0 = _mm256_loadu_ps(outptr); + _fsum1 = _mm256_loadu_ps(outptr + 8); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + __m256i _sum1 = _mm256_setzero_si256(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __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; + pB += 32; + } + if (max_kk0 >= 4) + { + __m128i _w_shift64 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _w_shift128 = _mm_unpacklo_epi64(_w_shift64, _w_shift64); + __m256i _w_shift0 = _mm256_broadcastsi128_si256(_w_shift128); + __m256i _w_shift1 = _mm256_shuffle_epi32(_w_shift0, _MM_SHUFFLE(2, 3, 0, 1)); + _sum0 = _mm256_sub_epi32(_sum0, _w_shift0); + _sum1 = _mm256_sub_epi32(_sum1, _w_shift1); + pA += 8; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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; + pB += 16; + } + for (; kk < max_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __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; + } + + __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; + } + + _mm256_storeu_ps(outptr, _fsum0); + _mm256_storeu_ps(outptr + 8, _fsum1); + outptr += 16; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)k * 4; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 4; + __m128 _fsum0; + __m128 _fsum1; + + if (k == 0) + { + _fsum0 = _mm_setzero_ps(); + _fsum1 = _mm_setzero_ps(); + } + else + { + _fsum0 = _mm_loadu_ps(outptr); + _fsum1 = _mm_loadu_ps(outptr + 4); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum0 = _mm_setzero_si128(); + __m128i _sum1 = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __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); +#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_kk0 >= 4) + { + __m128i _w_shift64 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _w_shift = _mm_unpacklo_epi64(_w_shift64, _w_shift64); + _sum0 = _mm_sub_epi32(_sum0, _w_shift); + _sum1 = _mm_sub_epi32(_sum1, _w_shift); + pA += 8; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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; + pB += 8; + } + for (; kk < max_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __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_cvtsi32_si128(((const int*)pB)[0]); + __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; + } + + __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; + } + + _mm_storeu_ps(outptr, _fsum0); + _mm_storeu_ps(outptr + 4, _fsum1); + outptr += 8; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // defined(__x86_64__) || defined(_M_X64) +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)k * 2; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 2; +#if __SSE2__ + __m128 _fsum; + + if (k == 0) + { + _fsum = _mm_setzero_ps(); + } + else + { + _fsum = _mm_loadu_ps(outptr); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __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__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 8; + pB += 8; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk0 >= 4) + { + __m128i _w_shift64 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _w_shift = _mm_unpacklo_epi32(_w_shift64, _w_shift64); + _sum = _mm_sub_epi32(_sum, _w_shift); + pA += 8; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_unpacklo_epi32(_pA16x1, _pA16x1); + __m128i _pB8 = _mm_cvtsi32_si128(((const int*)pB)[0]); + __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_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __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_cvtsi32_si128(((const unsigned short*)pB)[0]); + __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; + } + + __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; + } + + _mm_storeu_ps(outptr, _fsum); + outptr += 4; + +#else + float fsum00; + float fsum01; + float fsum10; + float fsum11; + + if (k == 0) + { + fsum00 = 0.f; + fsum01 = 0.f; + fsum10 = 0.f; + fsum11 = 0.f; + } + else + { + fsum00 = outptr[0]; + fsum01 = outptr[1]; + fsum10 = outptr[2]; + fsum11 = outptr[3]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + 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_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + int b0 = pB[0]; + int b1 = pB[1]; + sum00 += pA[0] * b0; + sum01 += pA[0] * b1; + sum10 += pA[1] * b0; + sum11 += pA[1] * b1; + b0 = pB[2]; + b1 = pB[3]; + sum00 += pA[2] * b0; + sum01 += pA[2] * b1; + sum10 += pA[3] * b0; + sum11 += pA[3] * b1; + b0 = pB[4]; + b1 = pB[5]; + sum00 += pA[4] * b0; + sum01 += pA[4] * b1; + sum10 += pA[5] * b0; + sum11 += pA[5] * b1; + b0 = pB[6]; + b1 = 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_kk0; kk++) + { + const int b0 = pB[0]; + const int b1 = 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__ + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + (size_t)k; + const float* pB_descales = pB_descales_panel + (size_t)block_start; +#if __SSE2__ + __m128 _fsum; + + if (k == 0) + { + _fsum = _mm_setzero_ps(); + } + else + { + _fsum = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)outptr); + } + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB8 = _mm_cvtsi32_si128(((const int*)pB)[0]); + __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_kk0 >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_loadl_epi64((const __m128i*)pA)); + pA += 8; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB8 = _mm_cvtsi32_si128(((const unsigned short*)pB)[0]); + __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_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __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(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++; + } + + __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; + float fsum1; + + if (k == 0) + { + fsum0 = 0.f; + fsum1 = 0.f; + } + else + { + fsum0 = outptr[0]; + fsum1 = outptr[1]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + int sum0 = 0; + int sum1 = 0; + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + int b0 = pB[0]; + sum0 += pA[0] * b0; + sum1 += pA[1] * b0; + b0 = pB[1]; + sum0 += pA[2] * b0; + sum1 += pA[3] * b0; + b0 = pB[2]; + sum0 += pA[4] * b0; + sum1 += pA[5] * b0; + b0 = pB[3]; + sum0 += pA[6] * b0; + sum1 += pA[7] * b0; + pA += 8; + pB += 4; + } + for (; kk < max_kk0; kk++) + { + const int b0 = 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__ + pB_panel += K; + pB_descales_panel += block_count; + } + + pAT += A_hstep * 2; + pAT_descales += A_descales_hstep * 2; + } + for (; ii < max_ii; ii++) + { + const signed char* pB_panel = pBT; + const float* pB_descales_panel = pBT_descales; + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB = pB_panel + (size_t)k * 8; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 8; + __m256 _fsum; + + if (k == 0) + { + _fsum = _mm256_setzero_ps(); + } + else + { + _fsum = _mm256_loadu_ps(outptr); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m256i _sum = _mm256_setzero_si256(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA32 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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_kk0 >= 4) + { + _sum = _mm256_sub_epi32(_sum, _mm256_set1_epi32(((const int*)pA)[0])); + pA += 4; + } +#else + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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; + pB += 32; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __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_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(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; + } + __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; + } + + _mm256_storeu_ps(outptr, _fsum); + outptr += 8; + pB_panel += (size_t)8 * K; + pB_descales_panel += (size_t)8 * block_count; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pB = pB_panel + (size_t)k * 4; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 4; + __m128 _fsum; + + if (k == 0) + { + _fsum = _mm_setzero_ps(); + } + else + { + _fsum = _mm_loadu_ps(outptr); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA32 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 4; + pB += 16; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk0 >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(((const int*)pA)[0])); + pA += 4; + } +#endif +#else + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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; + pB += 16; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __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_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(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_cvtsi32_si128(((const int*)pB)[0]); + __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; + } + __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; + } + + _mm_storeu_ps(outptr, _fsum); + outptr += 4; + pB_panel += (size_t)4 * K; + pB_descales_panel += (size_t)4 * block_count; + } +#endif // defined(__x86_64__) || defined(_M_X64) +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pB = pB_panel + (size_t)k * 2; + const float* pB_descales = pB_descales_panel + (size_t)block_start * 2; +#if __SSE2__ + __m128 _fsum; + + if (k == 0) + { + _fsum = _mm_setzero_ps(); + } + else + { + _fsum = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)outptr); + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA32 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 4; + pB += 8; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk0 >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(((const int*)pA)[0])); + pA += 4; + } +#endif +#else + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __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; + pB += 8; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __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_cvtsi32_si128(((const int*)pB)[0]); + __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_kk0; kk++) + { + __m128i _pA8 = _mm_cvtsi32_si128(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_cvtsi32_si128(((const unsigned short*)pB)[0]); + __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; + } + __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; + } + + _mm_storel_pi((__m64*)outptr, _fsum); + outptr += 2; +#else + float fsum0; + float fsum1; + + if (k == 0) + { + fsum0 = 0.f; + fsum1 = 0.f; + } + else + { + fsum0 = outptr[0]; + fsum1 = outptr[1]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + int sum0 = 0; + int sum1 = 0; + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + sum0 += pA[0] * pB[0]; + sum0 += pA[1] * pB[2]; + sum0 += pA[2] * pB[4]; + sum0 += pA[3] * pB[6]; + sum1 += pA[0] * pB[1]; + sum1 += pA[1] * pB[3]; + sum1 += pA[2] * pB[5]; + sum1 += pA[3] * pB[7]; + pA += 4; + pB += 8; + } + for (; kk < max_kk0; kk++) + { + sum0 += pA[0] * pB[0]; + sum1 += pA[0] * 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__ + pB_panel += (size_t)2 * K; + pB_descales_panel += (size_t)2 * block_count; + } + for (; jj < max_jj; jj++) + { + const signed char* pB = pB_panel + (size_t)k; + const float* pB_descales = pB_descales_panel + (size_t)block_start; +#if __SSE2__ + float fsum; + + if (k == 0) + { + fsum = 0.f; + } + else + { + fsum = outptr[0]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA = _mm_cvtsi32_si128(((const int*)pA)[0]); + __m128i _pB = _mm_cvtsi32_si128(((const int*)pB)[0]); +#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_kk0 >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_cvtsi32_si128(((const int*)pA)[0])); + pA += 4; + } +#endif +#else + for (; kk + 3 < max_kk0; kk += 4) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const int*)pA)[0]); + __m128i _pB8 = _mm_cvtsi32_si128(((const int*)pB)[0]); + __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; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk0; kk += 2) + { + __m128i _pA8 = _mm_cvtsi32_si128(((const unsigned short*)pA)[0]); + __m128i _pB8 = _mm_cvtsi32_si128(((const unsigned short*)pB)[0]); + __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_kk0; kk++) + { + __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++; + } + fsum += _mm_reduce_add_epi32(_sum) * pA_descales[0] * pB_descales[0]; + pA_descales += 1; + pB_descales++; + } + + outptr[0] = fsum; + outptr++; +#else + float fsum; + + if (k == 0) + { + fsum = 0.f; + } + else + { + fsum = outptr[0]; + } + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) + { + int sum = 0; + const int max_kk0 = std::min(tile_K - kk0, block_size); + int kk = 0; + for (; kk + 3 < max_kk0; kk += 4) + { + sum += pA[0] * pB[0]; + sum += pA[1] * pB[1]; + sum += pA[2] * pB[2]; + sum += pA[3] * pB[3]; + pA += 4; + pB += 4; + } + for (; kk < max_kk0; kk++) + { + sum += pA[0] * pB[0]; + pA++; + pB++; + } + + fsum += sum * pA_descales[0] * pB_descales[0]; + pA_descales += 1; + pB_descales++; + } + + outptr[0] = fsum; + outptr++; +#endif // __SSE2__ + pB_panel += K; + pB_descales_panel += block_count; + } + + 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, 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, 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, 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, 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, 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; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + + // topT microkernel lanes -> m0[n0..nNR-1], ..., mMR-1[n0..nNR-1] + 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) * 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((int)c_hstep)); + } + 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; + + // from + // 00 11 22 33 44 55 66 77 80 91 a2 b3 c4 d5 e6 f7 + // 01 12 23 30 45 56 67 74 81 92 a3 b0 c5 d6 e7 f4 + // 20 31 02 13 64 75 46 57 a0 b1 82 93 e4 f5 c6 d7 + // 21 32 03 10 65 76 47 54 a1 b2 83 90 e5 f6 c7 d4 + // 04 15 26 37 40 51 62 73 84 95 a6 b7 c0 d1 e2 f3 + // 05 16 27 34 41 52 63 70 85 96 a7 b4 c1 d2 e3 f0 + // 24 35 06 17 60 71 42 53 a4 b5 86 97 e0 f1 c2 d3 + // 25 36 07 14 61 72 43 50 a5 b6 87 94 e1 f2 c3 d0 + // + // to + // 00 10 20 30 40 50 60 70 80 90 a0 b0 c0 d0 e0 f0 + // 01 11 21 31 41 51 61 71 81 91 a1 b1 c1 d1 e1 f1 + // 02 12 22 32 42 52 62 72 82 92 a2 b2 c2 d2 e2 f2 + // 03 13 23 33 43 53 63 73 83 93 a3 b3 c3 d3 e3 f3 + // 04 14 24 34 44 54 64 74 84 94 a4 b4 c4 d4 e4 f4 + // 05 15 25 35 45 55 65 75 85 95 a5 b5 c5 d5 e5 f5 + // 06 16 26 36 46 56 66 76 86 96 a6 b6 c6 d6 e6 f6 + // 07 17 27 37 47 57 67 77 87 97 a7 b7 c7 d7 e7 f7 + _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) + { + __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) + { + __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) + { + __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + _f1 = _mm512_mul_ps(_f1, _alpha); + } + transpose16x2_ps(_f0, _f1); + { + __m128 _r = _mm512_extractf32x4_ps(_f0, 0); + _mm_storel_pi((__m64*)(p0), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep), _r); + } + { + __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); + } + { + __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); + } + { + __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); + } + { + __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); + } + { + __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); + } + { + __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); + } + { + __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) + { + __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) * c_hstep + 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); + pp += 64; + + // from + // 00 11 22 33 44 55 66 77 + // 01 12 23 30 45 56 67 74 + // 20 31 02 13 64 75 46 57 + // 21 32 03 10 65 76 47 54 + // 04 15 26 37 40 51 62 73 + // 05 16 27 34 41 52 63 70 + // 24 35 06 17 60 71 42 53 + // 25 36 07 14 61 72 43 50 + + // to + // 00 10 20 30 40 50 60 70 + // 01 11 21 31 41 51 61 71 + // 02 12 22 32 42 52 62 72 + // 03 13 23 33 43 53 63 73 + // 04 14 24 34 44 54 64 74 + // 05 15 25 35 45 55 65 75 + // 06 16 26 36 46 56 66 76 + // 07 17 27 37 47 57 67 77 + { + __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 + 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); + _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; + } +#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 +#if __AVX2__ + pp += 32; +#else + pp += 16; + pp1 += 16; +#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 + 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); + _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; + } +#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 +#if __AVX2__ + pp += 16; +#else + pp += 8; + pp1 += 8; +#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 + 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); + _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; + } + for (; jj < max_jj; jj++) + { +#if __AVX2__ + __m256 _f0 = _mm256_loadu_ps(pp); +#else + __m256 _f0 = _mm256_insertf128_ps(_mm256_castps128_ps256(_mm_loadu_ps(pp)), _mm_loadu_ps(pp1), 1); +#endif +#if __AVX2__ + pp += 8; +#else + pp += 4; + pp1 += 4; +#endif + if (pC) + { + if (broadcast_type_C == 0) + { + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); + } + 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) + { + __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) + _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 = 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) * c_hstep + 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); + pp += 32; + __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 + 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); + _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; + } +#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); + pp += 16; + { + _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 + 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); + _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; + } +#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); + pp += 8; + __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 + 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); + _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; + } + for (; jj < max_jj; jj += 1) + { + __m128 _f0 = _mm_loadu_ps(pp); + pp += 4; + if (pC) + { + if (broadcast_type_C == 0) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + } + 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) + { + __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) + _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++; + } + } + +#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) * c_hstep + 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); + pp += 16; + _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 + c_hstep); + 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; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = _mm_loadu_ps(pp + 0); + __m128 _f1 = _mm_loadu_ps(pp + 4); + pp += 8; + __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 + c_hstep); + 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; + } +#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)); + pp += 4; + 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 + c_hstep)); + 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; + } + for (; jj < max_jj; jj += 1) + { + __m128 _f0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pp); + pp += 2; + if (pC) + { + if (broadcast_type_C == 0) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _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) + _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++; + } +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + float f1_0 = pp[2]; + float f1_1 = pp[3]; + pp += 4; + 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) + { + 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; + } + } + 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; + + p0 += 2; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + float f1_0 = pp[1]; + pp += 2; + if (pC) + { + 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) + { + 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) + { + f0_0 *= alpha; + f1_0 *= alpha; + } + p0[0] = f0_0; + p0[out_hstep] = f1_0; + + p0++; + } + } + + 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) * 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) + 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); + pp += 8; + 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; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = _mm_loadu_ps(pp + 0); + pp += 4; + 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; + } +#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)); + pp += 2; + 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; + } +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + pp += 2; + 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; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + pp += 1; + 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++; + } + } +} + +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, 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, 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, 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, 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, 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; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + + // topT microkernel lanes -> n0[m0..mMR-1], ..., nNR-1[m0..mMR-1] + 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) * 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((int)c_hstep)); + } + 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; + + // from + // 00 11 22 33 44 55 66 77 80 91 a2 b3 c4 d5 e6 f7 + // 01 12 23 30 45 56 67 74 81 92 a3 b0 c5 d6 e7 f4 + // 20 31 02 13 64 75 46 57 a0 b1 82 93 e4 f5 c6 d7 + // 21 32 03 10 65 76 47 54 a1 b2 83 90 e5 f6 c7 d4 + // 04 15 26 37 40 51 62 73 84 95 a6 b7 c0 d1 e2 f3 + // 05 16 27 34 41 52 63 70 85 96 a7 b4 c1 d2 e3 f0 + // 24 35 06 17 60 71 42 53 a4 b5 86 97 e0 f1 c2 d3 + // 25 36 07 14 61 72 43 50 a5 b6 87 94 e1 f2 c3 d0 + // + // to + // 00 10 20 30 40 50 60 70 80 90 a0 b0 c0 d0 e0 f0 + // 01 11 21 31 41 51 61 71 81 91 a1 b1 c1 d1 e1 f1 + // 02 12 22 32 42 52 62 72 82 92 a2 b2 c2 d2 e2 f2 + // 03 13 23 33 43 53 63 73 83 93 a3 b3 c3 d3 e3 f3 + // 04 14 24 34 44 54 64 74 84 94 a4 b4 c4 d4 e4 f4 + // 05 15 25 35 45 55 65 75 85 95 a5 b5 c5 d5 e5 f5 + // 06 16 26 36 46 56 66 76 86 96 a6 b6 c6 d6 e6 f6 + // 07 17 27 37 47 57 67 77 87 97 a7 b7 c7 d7 e7 f7 + _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) + { + __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) + { + __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) + { + __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) + { + __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) * c_hstep + 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); + pp += 64; + + // from + // 00 11 22 33 44 55 66 77 + // 01 12 23 30 45 56 67 74 + // 20 31 02 13 64 75 46 57 + // 21 32 03 10 65 76 47 54 + // 04 15 26 37 40 51 62 73 + // 05 16 27 34 41 52 63 70 + // 24 35 06 17 60 71 42 53 + // 25 36 07 14 61 72 43 50 + + // to + // 00 10 20 30 40 50 60 70 + // 01 11 21 31 41 51 61 71 + // 02 12 22 32 42 52 62 72 + // 03 13 23 33 43 53 63 73 + // 04 14 24 34 44 54 64 74 + // 05 15 25 35 45 55 65 75 + // 06 16 26 36 46 56 66 76 + // 07 17 27 37 47 57 67 77 + { + __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 + 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); + _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; + } +#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 +#if __AVX2__ + pp += 32; +#else + pp += 16; + pp1 += 16; +#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 + 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); + _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; + } +#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 +#if __AVX2__ + pp += 16; +#else + pp += 8; + pp1 += 8; +#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 + 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); + _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; + } + 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 __AVX2__ + pp += 8; +#else + pp += 4; + pp1 += 4; +#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[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) + { + _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 = 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) * c_hstep + 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); + pp += 32; + __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 + 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); + _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; + } +#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); + pp += 16; + { + _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 + 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); + _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; + } +#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); + pp += 8; + __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 + 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); + _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; + } + for (; jj < max_jj; jj += 1) + { + __m128 _f = _mm_loadu_ps(pp); + pp += 4; + 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[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) + { + _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; + } + } + +#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) * c_hstep + 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); + pp += 16; + _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 + c_hstep); + 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; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _t0 = _mm_loadu_ps(pp + 0); + __m128 _t1 = _mm_loadu_ps(pp + 4); + pp += 8; + __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 + c_hstep); + 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; + } +#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)); + pp += 4; + 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 + c_hstep)); + 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; + } + for (; jj < max_jj; jj += 1) + { + __m128 _f = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pp); + pp += 2; + 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[c_hstep], 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; + } +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + float f1_0 = pp[2]; + float f1_1 = pp[3]; + pp += 4; + 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; + } + + 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[c_hstep] * beta; + f1_1 += pC[c_hstep + 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; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + float f1_0 = pp[1]; + pp += 2; + 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; + + 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[c_hstep] * 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; + } + } + + 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) * 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) + 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); + pp += 8; + 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; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f = _mm_loadu_ps(pp); + pp += 4; + 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; + } +#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); + pp += 2; + 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; + } +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + pp += 2; + 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; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + pp += 1; + 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; + } + } +} + +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(); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + + 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__ + 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 + + 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()); + + 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/K 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 + } + + 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 0d176d4778a..4ff8c1f393e 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,10 +7436,392 @@ 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 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; + + pack_B_tile_wq_int8(B, B_scales, packed_B, packed_B_descales, 0, N, K, block_size); + + return 0; +} + +int Gemm_x86::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; + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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, block_size, opt); + if (ret != 0) + return ret; + if (packed_B.empty() || packed_B_descales.empty()) + return -100; + + BT_data_wq_int8 = packed_B; + BT_data_wq_int8_descales = packed_B_descales; + B_data.release(); + B_data_quantize_scales.release(); + + 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, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + 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(); +#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, 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 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_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_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); + } + + 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 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()); + + 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, 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(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; + + #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 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()); + + 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) + { + 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); + + 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, 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; +} + +int Gemm_x86::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; + } + + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0 || weight_bits != 8) + return -1; + if (has_input_scale && B_data_input_scales.empty()) + return -100; + + 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) + { + if (output_N1M) + top_blob.create(M, 1, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + } + else + { + if (output_N1M) + top_blob.create(N, 1, M, (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, 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_x86::create_pipeline(const Option& opt) { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + if (weight_bits == 8) + return create_pipeline_wq_int8(opt); +#endif // NCNN_WEIGHT_QUANT + return 0; } @@ -7591,10 +7977,30 @@ int Gemm_x86::create_pipeline(const Option& opt) return 0; } +int Gemm_x86::destroy_pipeline(const Option& /*opt*/) +{ +#if NCNN_WEIGHT_QUANT + BT_data_wq_int8.release(); + BT_data_wq_int8_descales.release(); +#endif + + return 0; +} + 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 + int weight_bits; + int block_size; + bool has_input_scale; + if (get_weight_block_quantize_params(weight_bits, block_size, has_input_scale) != 0) + return -1; + if (weight_bits == 8) + return forward_wq_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 30263a951f6..d2b9476dffa 100644 --- a/src/layer/x86/gemm_x86.h +++ b/src/layer/x86/gemm_x86.h @@ -14,10 +14,15 @@ 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 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; @@ -32,6 +37,10 @@ class Gemm_x86 : 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/x86/gemm_x86_avx2.cpp b/src/layer/x86/gemm_x86_avx2.cpp index cab6c757975..3652c8023c9 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, Mat& BT_tile, Mat& BT_descales_tile, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, BT_tile, BT_descales_tile, 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 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, 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, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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 d8f46b5bba5..fe202e78bd4 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, Mat& BT_tile, Mat& BT_descales_tile, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, BT_tile, BT_descales_tile, 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 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, 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, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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 03120ff2e5b..2d2b624f0b3 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, Mat& BT_tile, Mat& BT_descales_tile, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, BT_tile, BT_descales_tile, 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 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, 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, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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 59e3f7e187b..6dfb061768c 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, Mat& BT_tile, Mat& BT_descales_tile, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, BT_tile, BT_descales_tile, 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 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, 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, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, 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 e2f58dbca5f..c82c76601a9 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 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, max_kk, 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 b1cc65423eb..f11eb3e601d 100644 --- a/src/layer/x86/multiheadattention_x86.cpp +++ b/src/layer/x86/multiheadattention_x86.cpp @@ -29,10 +29,362 @@ 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; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -45,18 +397,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 +449,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 +498,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) + { + destroy_pipeline(opt); + return ret; + } + ret = k_gemm->create_pipeline(opt); + if (ret != 0) { - k_weight_data.release(); - k_bias_data.release(); + 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 +547,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) + { + destroy_pipeline(opt); + return ret; + } + ret = v_gemm->create_pipeline(opt); + if (ret != 0) { - v_weight_data.release(); - v_bias_data.release(); + 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 +594,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 +653,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 +701,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; @@ -261,13 +743,33 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) int MultiHeadAttention_x86::destroy_pipeline(const Option& _opt) { if (weight_block_quantize) - return 0; + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 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 +780,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; } @@ -323,7 +825,17 @@ 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) - return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + { + int weight_bits; + int block_size; + bool has_input_scale; + const int ret = get_weight_block_quantize_params(weight_bits, block_size, has_input_scale); + if (ret != 0) + return ret; + + if (weight_bits != 8) + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + } int q_blob_i = 0; int k_blob_i = 0; @@ -341,10 +853,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 +911,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 +921,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 +949,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 +1000,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 +1028,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 +1075,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 66d88910c10..fb0fb6aa8ab 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; diff --git a/tests/test_gemm_block_quant.cpp b/tests/test_gemm_block_quant.cpp index 2d2b8cdd40d..c86b598d9e8 100644 --- a/tests/test_gemm_block_quant.cpp +++ b/tests/test_gemm_block_quant.cpp @@ -3,48 +3,33 @@ #include "testutil.h" -#include "gemm.h" -#include "layer_type.h" - -#include - -static void pack_signed_weight(unsigned char* ptr, int k, int bits, int q) +static int test_gemm_load_param(int quantize_term) { - const unsigned int mask = (1u << bits) - 1u; - const unsigned int v = (unsigned int)q & mask; - const int bit_offset = k * bits; + ncnn::ParamDict pd; + pd.set(3, 1); // transB + pd.set(5, 1); // constantB + pd.set(7, 1); + pd.set(8, 1); + pd.set(9, 32); + pd.set(18, quantize_term); - for (int b = 0; b < bits; b++) - { - if (v & (1u << b)) - { - const int out_bit = bit_offset + b; - ptr[out_bit / 8] |= (unsigned char)(1u << (out_bit % 8)); - } - } -} + ncnn::Layer* gemm = ncnn::create_layer_naive("Gemm"); + if (!gemm) + return -1; -static int float2int_weight(float v, int bits) -{ - const int qmax = (1 << (bits - 1)) - 1; - int q = (int)roundf(v); - if (q > qmax) q = qmax; - if (q < -qmax) q = -qmax; - return q; + const int ret = gemm->load_param(pd); + delete gemm; + return ret; } -static int weight_block_quantize_term(int bits, int block_size, int input_scale = 0) +#if NCNN_WEIGHT_QUANT +static int weight_block_quantize_term(int bits, int block_size, int has_input_scale = 0) { const int block_size_code = block_size == 32 ? 0 : block_size == 64 ? 1 : block_size == 128 ? 2 : -1; if ((bits != 4 && bits != 6 && bits != 8) || block_size_code < 0) return 0; - return bits * 100 + (input_scale ? 10 : 0) + block_size_code; -} - -static int weight_quantize_packed_k_bytes(int K, int bits) -{ - return (K * bits + 7) / 8; + return bits * 100 + (has_input_scale ? 10 : 0) + block_size_code; } static ncnn::Mat make_input_scales(int K) @@ -57,40 +42,69 @@ static ncnn::Mat make_input_scales(int K) return input_scales; } -static ncnn::Mat scale_weight_by_input_scales(const ncnn::Mat& weight_data, const ncnn::Mat& input_scales, int inverse) +static void RandomizeA(ncnn::Mat& A, int transA, int block_size, const ncnn::Mat& input_scales) { - const int K = weight_data.w; - const int N = weight_data.h; + int M = A.h; + if (transA) + M = A.w; + if (A.dims == 3) + M = A.c; - ncnn::Mat weight_data1(K, N); + const int K = transA ? A.h : A.w; const float* input_scale_ptr = input_scales; - for (int n = 0; n < N; n++) + for (int i = 0; i < M; i++) { - const float* ptr = weight_data.row(n); - float* outptr = weight_data1.row(n); + float* ptr = 0; + if (!transA) + ptr = A.dims == 3 ? A.channel(i) : A.row(i); for (int k = 0; k < K; k++) - outptr[k] = inverse ? ptr[k] / input_scale_ptr[k] : ptr[k] * input_scale_ptr[k]; + { + int q = RandomInt(-120, 121); + if (k % block_size == 0) + q = (i + k / block_size) % 2 == 0 ? 127 : -127; + + float v = q / 64.f; + if (input_scale_ptr) + v /= input_scale_ptr[k]; + + if (transA) + A.row(k)[i] = v; + else + ptr[k] = v; + } } +} - return weight_data1; +static void pack_signed_weight(unsigned char* ptr, int k, int bits, int q) +{ + const unsigned int v = (unsigned int)q & ((1u << bits) - 1u); + const int bit_offset = k * bits; + + for (int b = 0; b < bits; b++) + { + if (v & (1u << b)) + { + const int out_bit = bit_offset + b; + ptr[out_bit / 8] |= (unsigned char)(1u << (out_bit % 8)); + } + } } static int quantize_weight(const ncnn::Mat& weight_data, int bits, int block_size, ncnn::Mat& weight_data_quantized, ncnn::Mat& weight_data_quantize_scales, ncnn::Mat& weight_data_dequantized) { const int K = weight_data.w; const int N = weight_data.h; - const int packed_k_bytes = weight_quantize_packed_k_bytes(K, bits); const int block_count = (K + block_size - 1) / block_size; - weight_data_quantized.create(packed_k_bytes, N, (size_t)1u); + weight_data_quantized.create((K * bits + 7) / 8, N, (size_t)1u); weight_data_quantize_scales.create(block_count, N); weight_data_dequantized.create(K, N); if (weight_data_quantized.empty() || weight_data_quantize_scales.empty() || weight_data_dequantized.empty()) return -100; - memset(weight_data_quantized.data, 0, weight_data_quantized.total() * weight_data_quantized.elemsize); + weight_data_quantized.fill(0); const int qmax = (1 << (bits - 1)) - 1; for (int n = 0; n < N; n++) @@ -118,7 +132,9 @@ static int quantize_weight(const ncnn::Mat& weight_data, int bits, int block_siz for (int k = 0; k < max_kk; k++) { - const int q = float2int_weight(ptr[k0 + k] * scale, bits); + int q = (int)roundf(ptr[k0 + k] * scale); + if (q > qmax) q = qmax; + if (q < -qmax) q = -qmax; pack_signed_weight(qptr, k0 + k, bits, q); deqptr[k0 + k] = q / scale; } @@ -128,12 +144,12 @@ static int quantize_weight(const ncnn::Mat& weight_data, int bits, int block_siz return 0; } -static ncnn::ParamDict make_gemm_param(int M, int N, int K, int quantize_term, float alpha = 1.f, float beta = 1.f, int constantC = 0, int broadcast_type_C = 0) +static ncnn::ParamDict make_gemm_param(int M, int N, int K, int quantize_term, float alpha = 1.f, float beta = 1.f, int constantC = 0, int broadcast_type_C = -1, int transA = 0, int output_transpose = 0, int output_N1M = 0) { ncnn::ParamDict pd; pd.set(0, alpha); pd.set(1, beta); - pd.set(2, 0); // transA + pd.set(2, transA); pd.set(3, 1); // transB pd.set(4, 0); // constantA pd.set(5, 1); // constantB @@ -142,105 +158,487 @@ static ncnn::ParamDict make_gemm_param(int M, int N, int K, int quantize_term, f pd.set(8, N); pd.set(9, K); pd.set(10, broadcast_type_C); - if (quantize_term) - pd.set(18, quantize_term); + pd.set(11, output_N1M); + pd.set(14, output_transpose); + pd.set(18, quantize_term); return pd; } -static int test_gemm_block_quant(int M, int N, int K, int bits, int block_size, int broadcast_type_C, int input_scale = 0, float alpha = 1.f, float beta = 1.f, int constantC = 0, int dims3 = 0) +static int test_gemm_invalid_weight_block_quantize_term() { - const int quantize_term = weight_block_quantize_term(bits, block_size, input_scale); + const int invalid_quantize_terms[] = {403, 420, 700}; - ncnn::Mat A = dims3 ? RandomMat(K, 1, M, -2.f, 2.f) : RandomMat(K, M, -2.f, 2.f); - ncnn::Mat B = RandomMat(K, N, -2.f, 2.f); - ncnn::Mat C; + for (int i = 0; i < (int)(sizeof(invalid_quantize_terms) / sizeof(invalid_quantize_terms[0])); i++) + { + if (test_gemm_load_param(invalid_quantize_terms[i]) == 0) + { + fprintf(stderr, "test_gemm_invalid_weight_block_quantize_term accepted quantize_term=%d\n", invalid_quantize_terms[i]); + return -1; + } + } + + return 0; +} + +static int test_gemm_block_quant(const ncnn::Mat& A, const ncnn::Mat& B, const ncnn::Mat& C, const ncnn::ParamDict& pd, int bits, int block_size, int has_input_scale, int constantC) +{ + const int K = B.w; + const int N = B.h; ncnn::Mat B_input_scales; - if (input_scale) + ncnn::Mat B1 = B; + if (has_input_scale) { B_input_scales = make_input_scales(K); - B = scale_weight_by_input_scales(B, B_input_scales, 1); + B1 = B.clone(); + + const float* input_scale_ptr = B_input_scales; + for (int n = 0; n < N; n++) + { + float* ptr = B1.row(n); + for (int k = 0; k < K; k++) + ptr[k] /= input_scale_ptr[k]; + } } ncnn::Mat B_quantized; ncnn::Mat B_quantize_scales; ncnn::Mat B_dequantized; - int ret = quantize_weight(B, bits, block_size, B_quantized, B_quantize_scales, B_dequantized); + int ret = quantize_weight(B1, bits, block_size, B_quantized, B_quantize_scales, B_dequantized); if (ret != 0) return ret; - if (input_scale) - B_dequantized = scale_weight_by_input_scales(B_dequantized, B_input_scales, 0); - - if (broadcast_type_C == 0) C = RandomMat(1); - if (broadcast_type_C == 1) C = RandomMat(M); - if (broadcast_type_C == 2) C = RandomMat(1, M); - if (broadcast_type_C == 3) C = RandomMat(N, M); - if (broadcast_type_C == 4) C = RandomMat(N); + if (has_input_scale) + { + const float* input_scale_ptr = B_input_scales; + for (int n = 0; n < N; n++) + { + float* ptr = B_dequantized.row(n); + for (int k = 0; k < K; k++) + ptr[k] *= input_scale_ptr[k]; + } + } std::vector weights; weights.push_back(B_quantized); if (constantC) weights.push_back(C); weights.push_back(B_quantize_scales); - if (input_scale) + if (has_input_scale) weights.push_back(B_input_scales); - std::vector ref_weights; - ref_weights.push_back(B_dequantized); - if (constantC) - ref_weights.push_back(C); - - std::vector inputs; - inputs.push_back(A); - + std::vector a; + a.push_back(A); if (!constantC && !C.empty()) - inputs.push_back(C); + a.push_back(C); ncnn::Option opt; - opt.num_threads = 2; opt.use_packing_layout = false; opt.use_fp16_packed = false; opt.use_fp16_storage = false; opt.use_fp16_arithmetic = false; opt.use_bf16_storage = false; - std::vector outputs; - std::vector refs; - ret = test_layer_cpu(ncnn::LayerType::Gemm, make_gemm_param(M, N, K, quantize_term, alpha, beta, constantC, broadcast_type_C), weights, opt, inputs, 1, outputs, std::vector(), TEST_LAYER_DISABLE_GPU_TESTING); - if (ret != 0) + if (bits != 8) { - fprintf(stderr, "test_gemm_block_quant failed ret=%d M=%d N=%d K=%d bits=%d block_size=%d broadcast_type_C=%d input_scale=%d constantC=%d dims3=%d\n", ret, M, N, K, bits, block_size, broadcast_type_C, input_scale, constantC, dims3); + std::vector ref_weights; + ref_weights.push_back(B_dequantized); + if (constantC) + ref_weights.push_back(C); + + ncnn::ParamDict ref_pd = pd; + ref_pd.set(18, 0); + + std::vector refs; + ret = test_layer_naive(ncnn::layer_to_index("Gemm"), ref_pd, ref_weights, a, 1, refs, TEST_LAYER_DISABLE_GPU_TESTING); + if (ret != 0) + return ret; + + for (int t = 0; t < 2; t++) + { + std::vector outputs; + const int flag = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_cpu(ncnn::layer_to_index("Gemm"), pd, weights, opt, a, 1, outputs, std::vector(), flag); + if (ret != 0 || CompareMat(outputs, refs, 0.001f) != 0) + return ret != 0 ? ret : -1; + } + + return 0; + } + + ret = test_layer_opt("Gemm", pd, weights, opt, a, 1, 0.001f, TEST_LAYER_DISABLE_GPU_TESTING); + if (ret != 0) return ret; + + return test_layer_opt("Gemm", pd, weights, opt, a, 1, 0.001f, TEST_LAYER_DISABLE_GPU_TESTING | TEST_LAYER_ENABLE_THREADING); +} + +static int test_gemm(int M, int N, int K, int bits, int block_size, int has_input_scale = 0, int transA = 0, int output_transpose = 0, int output_N1M = 0, int dims3 = 0) +{ + ncnn::Mat A = dims3 || output_N1M ? ncnn::Mat(K, 1, M) : transA ? ncnn::Mat(M, K) : ncnn::Mat(K, M); + ncnn::Mat B(K, N); + if (bits == 8) + RandomizeA(A, transA, block_size, has_input_scale ? make_input_scales(K) : ncnn::Mat()); + else + Randomize(A, -2.f, 2.f); + Randomize(B, -2.f, 2.f); + + const ncnn::ParamDict pd = make_gemm_param(M, N, K, weight_block_quantize_term(bits, block_size, has_input_scale), 1.f, 1.f, 0, -1, transA, output_transpose, output_N1M); + const int ret = test_gemm_block_quant(A, B, ncnn::Mat(), pd, bits, block_size, has_input_scale, 0); + if (ret != 0) + { + fprintf(stderr, "test_gemm failed M=%d N=%d K=%d bits=%d block_size=%d has_input_scale=%d transA=%d output_transpose=%d output_N1M=%d dims3=%d\n", M, N, K, bits, block_size, has_input_scale, transA, output_transpose, output_N1M, dims3); } - ret = test_layer_cpu(ncnn::LayerType::Gemm, make_gemm_param(M, N, K, 0, alpha, beta, constantC, broadcast_type_C), ref_weights, opt, inputs, 1, refs, std::vector(), TEST_LAYER_DISABLE_GPU_TESTING); + return ret; +} + +static int test_gemm_bias(int M, int N, int K, int bits, int block_size, const ncnn::Mat& C, float alpha, float beta, int has_input_scale, int transA, int output_transpose, int constantC, int output_N1M = 0) +{ + int broadcast_type_C = 4; + if (C.dims == 1 && C.w == 1) + broadcast_type_C = 0; + else if (C.dims == 1 && C.w == M) + broadcast_type_C = 1; + else if (C.dims == 2 && C.w == 1 && C.h == M) + broadcast_type_C = 2; + else if (C.dims == 2 && C.w == N && C.h == M) + broadcast_type_C = 3; + + ncnn::Mat A = output_N1M ? ncnn::Mat(K, 1, M) : transA ? ncnn::Mat(M, K) : ncnn::Mat(K, M); + ncnn::Mat B(K, N); + if (bits == 8) + RandomizeA(A, transA, block_size, has_input_scale ? make_input_scales(K) : ncnn::Mat()); + else + Randomize(A, -2.f, 2.f); + Randomize(B, -2.f, 2.f); + + const ncnn::ParamDict pd = make_gemm_param(M, N, K, weight_block_quantize_term(bits, block_size, has_input_scale), alpha, beta, constantC, broadcast_type_C, transA, output_transpose, output_N1M); + const int ret = test_gemm_block_quant(A, B, C, pd, bits, block_size, has_input_scale, constantC); if (ret != 0) { - fprintf(stderr, "test_gemm_block_quant reference failed ret=%d M=%d N=%d K=%d\n", ret, M, N, K); - return ret; + fprintf(stderr, "test_gemm_bias failed M=%d N=%d K=%d bits=%d block_size=%d C.dims=%d C=(%d %d) alpha=%f beta=%f has_input_scale=%d transA=%d output_transpose=%d constantC=%d output_N1M=%d\n", M, N, K, bits, block_size, C.dims, C.w, C.h, alpha, beta, has_input_scale, transA, output_transpose, constantC, output_N1M); + } + + return ret; +} + +static int test_gemm_w8a8_zero(int zero_A, int zero_B_block) +{ + const int M = 31; + const int N = 16; + const int K = 67; + const int block_size = 32; + + ncnn::Mat A(K, M); + ncnn::Mat B(K, N); + RandomizeA(A, 0, block_size, ncnn::Mat()); + Randomize(B, -2.f, 2.f); + + if (zero_A) + A.fill(0.f); + if (zero_B_block) + { + for (int n = 0; n < N; n++) + { + float* ptr = B.row(n); + for (int k = 0; k < block_size; k++) + ptr[k] = 0.f; + } + } + + const ncnn::ParamDict pd = make_gemm_param(M, N, K, weight_block_quantize_term(8, block_size)); + const int ret = test_gemm_block_quant(A, B, ncnn::Mat(), pd, 8, block_size, 0, 0); + if (ret != 0) + fprintf(stderr, "test_gemm_w8a8_zero failed zero_A=%d zero_B_block=%d\n", zero_A, zero_B_block); + + return ret; +} + +static int test_gemm_w8a8_tile(int M, int N, int K, int block_size, int TILE_M, int TILE_N, int TILE_K) +{ + ncnn::Mat A(K, M); + ncnn::Mat B(K, N); + RandomizeA(A, 0, block_size, ncnn::Mat()); + Randomize(B, -2.f, 2.f); + + ncnn::ParamDict pd = make_gemm_param(M, N, K, weight_block_quantize_term(8, block_size)); + pd.set(20, TILE_M); + pd.set(21, TILE_N); + pd.set(22, TILE_K); + + const int ret = test_gemm_block_quant(A, B, ncnn::Mat(), pd, 8, block_size, 0, 0); + if (ret != 0) + fprintf(stderr, "test_gemm_w8a8_tile failed M=%d N=%d K=%d TILE_M=%d TILE_N=%d TILE_K=%d\n", M, N, K, TILE_M, TILE_N, TILE_K); + + return ret; +} + +static int float2int8_reference(float v) +{ + int q = (int)roundf(v); + if (q > 127) return 127; + if (q < -127) return -127; + return q; +} + +static void reference_gemm_w8a8(const ncnn::Mat& A, const ncnn::Mat& B, const ncnn::Mat& B_scales, const ncnn::Mat& input_scales, int block_size, ncnn::Mat& reference) +{ + const int M = A.h; + const int N = B.h; + const int K = A.w; + const int block_count = (K + block_size - 1) / block_size; + const float* input_scale_ptr = input_scales; + + reference.create(N, M); + for (int i = 0; i < M; i++) + { + const float* ptrA = A.row(i); + float* outptr = reference.row(i); + + for (int j = 0; j < N; j++) + { + const signed char* ptrB = B.row(j); + const float* scale_ptr = B_scales.row(j); + float sum = 0.f; + + 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 float v = fabsf(ptrA[k0 + kk] * input_scale_ptr[k0 + kk]); + if (v > absmax) + absmax = v; + } + + if (absmax == 0.f) + continue; + + const float scale = 127.f / absmax; + int sum_int32 = 0; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + sum_int32 += float2int8_reference(ptrA[k] * input_scale_ptr[k] * scale) * ptrB[k]; + } + + sum += sum_int32 * (absmax / 127.f) / scale_ptr[g]; + } + + outptr[j] = sum; + } + } +} + +static int test_gemm_w8a8_quantize_rounding(int has_input_scale) +{ + const int M = 1; + const int N = 2; + const int K = 64; + const int block_size = 32; + + ncnn::Mat A(K, M); + A.fill(0.f); + A[0] = 0.25f; + A[1] = -0.25f; + A[2] = 126.25f; + A[3] = -126.25f; + A[4] = 127.f; + A[5] = -127.f; + + ncnn::Mat B_input_scales; + if (has_input_scale) + { + B_input_scales.create(K); + B_input_scales.fill(2.f); + for (int k = 0; k < K; k++) + A[k] *= 0.5f; + } + + ncnn::Mat B_quantized(K, N, (size_t)1u); + B_quantized.fill(7); + signed char* ptrB0 = B_quantized.row(0); + signed char* ptrB1 = B_quantized.row(1); + for (int k = 0; k < 6; k++) + { + ptrB0[k] = (signed char)(k + 1); + ptrB1[k] = (signed char)(6 - k); + } + + ncnn::Mat B_quantize_scales(2, N); + B_quantize_scales.fill(1.f); + + ncnn::Mat reference(N, M); + reference[0] = -253.f; + reference[1] = 253.f; + + std::vector weights; + weights.push_back(B_quantized); + weights.push_back(B_quantize_scales); + if (has_input_scale) + weights.push_back(B_input_scales); + + std::vector a(1, A); + ncnn::Option opt; + opt.use_packing_layout = false; + opt.use_fp16_packed = false; + opt.use_fp16_storage = false; + opt.use_fp16_arithmetic = false; + opt.use_bf16_storage = false; + + const ncnn::ParamDict pd = make_gemm_param(M, N, K, weight_block_quantize_term(8, block_size, has_input_scale)); + int ret = 0; + for (int t = 0; t < 2; t++) + { + std::vector outputs; + const int flag = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_cpu(ncnn::layer_to_index("Gemm"), pd, weights, opt, a, 1, outputs, std::vector(), flag); + if (ret != 0 || CompareMat(outputs[0], reference, 0.f) != 0) + { + fprintf(stderr, "test_gemm_w8a8_quantize_rounding failed has_input_scale=%d threading=%d\n", has_input_scale, t); + return ret != 0 ? ret : -1; + } } - ret = CompareMat(outputs, refs, 0.001f); + return 0; +} + +static int test_gemm_w8a8_pipeline_reuse() +{ + const int N = 31; + const int K = 67; + const int block_size = 32; + + ncnn::Mat B_input_scales = make_input_scales(K); + ncnn::Mat B_quantized(K, N, (size_t)1u); + ncnn::Mat B_quantize_scales(3, N); + B_quantize_scales.fill(1.f); + for (int n = 0; n < N; n++) + { + signed char* ptr = B_quantized.row(n); + for (int k = 0; k < K; k++) + ptr[k] = (signed char)((n * 19 + k * 11) % 101 - 50); + } + + std::vector weights(3); + weights[0] = B_quantized; + weights[1] = B_quantize_scales; + weights[2] = B_input_scales; + + ncnn::Layer* gemm = ncnn::create_layer_cpu("Gemm"); + if (!gemm) + return -1; + + int ret = gemm->load_param(make_gemm_param(0, N, K, weight_block_quantize_term(8, block_size, 1))); + if (ret == 0) + ret = gemm->load_model(ncnn::ModelBinFromMatArray(weights.data())); + + ncnn::Option opt; + opt.num_threads = 2; + opt.use_packing_layout = false; + opt.use_fp16_packed = false; + opt.use_fp16_storage = false; + opt.use_fp16_arithmetic = false; + opt.use_bf16_storage = false; + + if (ret == 0) + ret = gemm->create_pipeline(opt); if (ret != 0) { - fprintf(stderr, "test_gemm_block_quant compare failed M=%d N=%d K=%d bits=%d block_size=%d broadcast_type_C=%d input_scale=%d constantC=%d dims3=%d\n", M, N, K, bits, block_size, broadcast_type_C, input_scale, constantC, dims3); + delete gemm; return ret; } - return 0; + const int test_M[] = {128, 1, 7}; + const int test_threads[] = {2, 1, 4}; + for (int t = 0; t < 3; t++) + { + const int M = test_M[t]; + ncnn::Mat A(K, M); + RandomizeA(A, 0, block_size, B_input_scales); + + ncnn::Mat reference; + reference_gemm_w8a8(A, B_quantized, B_quantize_scales, B_input_scales, block_size, reference); + + std::vector bottom_blobs(1, A); + std::vector top_blobs(1); + ncnn::Option opt1 = opt; + opt1.num_threads = test_threads[t]; + ret = gemm->forward(bottom_blobs, top_blobs, opt1); + if (ret != 0 || CompareMat(top_blobs[0], reference, 0.001f) != 0) + { + fprintf(stderr, "test_gemm_w8a8_pipeline_reuse failed M=%d threads=%d\n", M, opt1.num_threads); + ret = ret != 0 ? ret : -1; + break; + } + } + + const int destroy_ret = gemm->destroy_pipeline(opt); + delete gemm; + return ret != 0 ? ret : destroy_ret; } +static int test_gemm_0() +{ + return 0 + || test_gemm(3, 5, 65, 4, 64) + || test_gemm(3, 4, 33, 4, 32, 1) + || test_gemm_bias(3, 4, 33, 4, 32, RandomMat(1), 1.7f, 0.3f, 0, 0, 0, 1) + || test_gemm(4, 7, 67, 6, 64) + || test_gemm(2, 4, 31, 6, 32, 1) + || test_gemm(3, 4, 129, 6, 128) + || test_gemm_bias(4, 7, 67, 6, 64, RandomMat(7, 4), 0.7f, 1.3f, 1, 0, 0, 0); +} + +static int test_gemm_1(int M, int N, int K, int block_size) +{ + return 0 + || test_gemm(M, N, K, 8, block_size) + || test_gemm(M, N, K, 8, block_size, 1, 0, 1) + || test_gemm(M, N, K, 8, block_size, 0, 1); +} + +static int test_gemm_2() +{ + const int M = 5; + const int N = 7; + const int K = 35; + const int block_size = 32; + + return 0 + || test_gemm_bias(M, N, K, 8, block_size, RandomMat(1), 0.7f, 1.3f, 0, 0, 0, 0) + || test_gemm_bias(M, N, K, 8, block_size, RandomMat(M), 1.f, 0.3f, 1, 0, 0, 1) + || test_gemm_bias(M, N, K, 8, block_size, RandomMat(1, M), 0.7f, 1.f, 0, 0, 1, 0) + || test_gemm_bias(M, N, K, 8, block_size, RandomMat(N, M), 1.7f, 0.3f, 1, 0, 1, 1) + || test_gemm_bias(M, N, K, 8, block_size, RandomMat(N), 0.7f, 1.3f, 0, 1, 0, 0) + || test_gemm_bias(M, N, K, 8, block_size, RandomMat(N), 1.f, 0.f, 0, 0, 0, 0) + || test_gemm(3, 5, 67, 8, 64, 0, 0, 0, 1) + || test_gemm_bias(3, 5, 67, 8, 64, RandomMat(5, 3), 0.7f, 0.3f, 1, 0, 1, 0, 1); +} + +static int test_gemm_3() +{ + return 0 + || test_gemm_w8a8_zero(1, 0) + || test_gemm_w8a8_zero(0, 1) + || test_gemm(5, 3, 65, 8, 64, 0, 0, 0, 0, 1) + || test_gemm(7, 9, 67, 8, 64, 1, 1, 1) + || test_gemm_w8a8_tile(13, 19, 35, 32, 7, 5, 17) + || test_gemm_w8a8_tile(7, 9, 67, 64, 5, 7, 48); +} +#endif // NCNN_WEIGHT_QUANT + int main() { SRAND(7767517); #if !NCNN_WEIGHT_QUANT - ncnn::ParamDict pd = make_gemm_param(3, 4, 5, 410); - - ncnn::Gemm gemm; - if (gemm.load_param(pd) == 0) + if (test_gemm_load_param(410) == 0) { fprintf(stderr, "test_gemm_block_quant failed NCNN_WEIGHT_QUANT=OFF accepted weight block quantization\n"); return -1; @@ -248,18 +646,36 @@ int main() return 0; #else - return 0 - || test_gemm_block_quant(3, 5, 65, 4, 64, -1) - || test_gemm_block_quant(3, 4, 33, 4, 32, 0) - || test_gemm_block_quant(4, 7, 67, 6, 64, 3) - || test_gemm_block_quant(3, 4, 33, 6, 32, 2) - || test_gemm_block_quant(2, 3, 5, 8, 32, 4) - || test_gemm_block_quant(2, 5, 65, 8, 32, -1) - || test_gemm_block_quant(3, 4, 129, 6, 128, -1) - || test_gemm_block_quant(3, 5, 65, 4, 64, -1, 1) - || test_gemm_block_quant(2, 4, 31, 6, 32, 1, 1) - || test_gemm_block_quant(3, 4, 32, 4, 32, -1) - || test_gemm_block_quant(3, 4, 33, 4, 32, 4, 0, 1.7f, 0.3f, 1) - || test_gemm_block_quant(5, 3, 65, 6, 64, -1, 0, 1.f, 1.f, 0, 1); + int ret = test_gemm_invalid_weight_block_quantize_term() + || test_gemm_w8a8_quantize_rounding(0) + || test_gemm_w8a8_quantize_rounding(1) + || test_gemm_w8a8_pipeline_reuse() + || test_gemm_0() + || test_gemm_2() + || test_gemm_3(); + if (ret != 0) + return ret; + + const int mnkb[][4] = { + {1, 7, 31, 32}, + {2, 3, 32, 32}, + {3, 6, 33, 32}, + {8, 3, 34, 32}, + {16, 16, 35, 32}, + {17, 31, 64, 32}, + {31, 17, 65, 64}, + {3, 129, 67, 64}, + {8, 16, 128, 128}, + {19, 3, 129, 128} + }; + + for (int i = 0; i < (int)(sizeof(mnkb) / sizeof(mnkb[0])); i++) + { + ret = test_gemm_1(mnkb[i][0], mnkb[i][1], mnkb[i][2], mnkb[i][3]); + if (ret != 0) + return ret; + } + + return 0; #endif } diff --git a/tests/test_gemm_oom.cpp b/tests/test_gemm_oom.cpp index b771144e493..448ad5fa2ce 100644 --- a/tests/test_gemm_oom.cpp +++ b/tests/test_gemm_oom.cpp @@ -149,6 +149,45 @@ static int test_gemm_1(int M, int N, int K) || test_gemm_bias_oom(M, N, K, RandomMat(N, M), 5.1f, 0.8f, 1, 1, 1, 1, 1, 1); } +#if NCNN_WEIGHT_QUANT +static int test_gemm_w8a8_oom(int M, int N, int K, int block_size, int input_scale, int output_transpose) +{ + const int block_count = (K + block_size - 1) / block_size; + const int block_size_code = block_size == 32 ? 0 : block_size == 64 ? 1 : 2; + + ncnn::ParamDict pd; + pd.set(0, 1.7f); + pd.set(1, 0.3f); + pd.set(2, 0); + pd.set(3, 1); + pd.set(4, 0); + pd.set(5, 1); + pd.set(6, 1); + pd.set(7, M); + pd.set(8, N); + pd.set(9, K); + pd.set(10, 4); + pd.set(14, output_transpose); + pd.set(18, 800 + input_scale * 10 + block_size_code); + + std::vector weights; + weights.push_back(RandomS8Mat(K, N)); + weights.push_back(RandomMat(N)); + weights.push_back(RandomMat(block_count, N, 10.f, 20.f)); + if (input_scale) + weights.push_back(RandomMat(K, 0.5f, 1.5f)); + + const ncnn::Mat A = RandomMat(K, M); + int ret = test_layer_oom("Gemm", pd, weights, A, TEST_LAYER_ENABLE_THREADING); + if (ret != 0) + { + fprintf(stderr, "test_gemm_w8a8_oom failed M=%d N=%d K=%d block_size=%d input_scale=%d output_transpose=%d\n", M, N, K, block_size, input_scale, output_transpose); + } + + return ret; +} +#endif // NCNN_WEIGHT_QUANT + #if NCNN_INT8 static int test_gemm_int8_oom(int M, int N, int K, int transA, int transB, int output_elemtype, int output_transpose, int constantA, int constantB, int output_N1M) { @@ -462,5 +501,11 @@ int main() return ret3; } +#if NCNN_WEIGHT_QUANT + return 0 + || test_gemm_w8a8_oom(3, 5, 65, 32, 0, 0) + || test_gemm_w8a8_oom(8, 17, 129, 128, 1, 1); +#else return 0; +#endif } diff --git a/tests/test_multiheadattention_block_quant.cpp b/tests/test_multiheadattention_block_quant.cpp index e059bd486c3..68bf9bd3062 100644 --- a/tests/test_multiheadattention_block_quant.cpp +++ b/tests/test_multiheadattention_block_quant.cpp @@ -3,15 +3,10 @@ #include "testutil.h" -#include "layer_type.h" -#include "multiheadattention.h" - -#include - +#if NCNN_WEIGHT_QUANT static void pack_signed_weight(unsigned char* ptr, int k, int bits, int q) { - const unsigned int mask = (1u << bits) - 1u; - const unsigned int v = (unsigned int)q & mask; + const unsigned int v = (unsigned int)q & ((1u << bits) - 1); const int bit_offset = k * bits; for (int b = 0; b < bits; b++) @@ -33,71 +28,82 @@ static int float2int_weight(float v, int bits) return q; } -static int weight_block_quantize_term(int bits, int block_size, int input_scale = 0) +static int weight_block_quantize_term(int bits, int block_size, int has_input_scale) { - const int block_size_code = block_size == 32 ? 0 : block_size == 64 ? 1 : block_size == 128 ? 2 : -1; - if ((bits != 4 && bits != 6 && bits != 8) || block_size_code < 0) - return 0; - - return bits * 100 + (input_scale ? 10 : 0) + block_size_code; + const int block_size_code = block_size == 32 ? 0 : block_size == 64 ? 1 : 2; + return bits * 100 + (has_input_scale ? 10 : 0) + block_size_code; } -static int weight_quantize_packed_k_bytes(int K, int bits) +static ncnn::Mat make_input_scales(int size, int offset) { - return (K * bits + 7) / 8; + ncnn::Mat scales(size); + float* ptr = scales; + const float scale_table[5] = {0.5f, 1.f, 2.f, 0.25f, 4.f}; + for (int i = 0; i < size; i++) + ptr[i] = scale_table[(i + offset) % 5]; + + return scales; } -static ncnn::Mat make_input_scales(int K) +static ncnn::Mat make_w8a8_mat(int width, int height, int block_size, const ncnn::Mat& input_scales = ncnn::Mat()) { - ncnn::Mat input_scales(K); - float* ptr = input_scales; - for (int k = 0; k < K; k++) - ptr[k] = 0.75f + (k % 5) * 0.15f; + ncnn::Mat m(width, height); + const float* scale_ptr = input_scales; + + for (int y = 0; y < height; y++) + { + float* ptr = m.row(y); + for (int x = 0; x < width; x++) + { + int q = RandomInt(-120, 121); + if (x % block_size == 0) + q = y % 2 == 0 ? 127 : -127; - return input_scales; + ptr[x] = scale_ptr ? q / (64.f * scale_ptr[x]) : q / 64.f; + } + } + + return m; } -static ncnn::Mat scale_weight_by_input_scales(const ncnn::Mat& weight_data, const ncnn::Mat& input_scales, int inverse) +static ncnn::Mat make_w8a8_cache(int width, int height) { - const int K = weight_data.w; - const int N = weight_data.h; + ncnn::Mat m(width, height); + std::vector values(width); + for (int x = 0; x < width; x++) + values[x] = RandomInt(-120, 121) / 64.f; - ncnn::Mat weight_data1(K, N); - const float* input_scale_ptr = input_scales; - - for (int n = 0; n < N; n++) + for (int y = 0; y < height; y++) { - const float* ptr = weight_data.row(n); - float* outptr = weight_data1.row(n); - - for (int k = 0; k < K; k++) - outptr[k] = inverse ? ptr[k] / input_scale_ptr[k] : ptr[k] * input_scale_ptr[k]; + float* ptr = m.row(y); + for (int x = 0; x < width; x++) + ptr[x] = values[x]; } - return weight_data1; + return m; } -static int quantize_weight(const ncnn::Mat& weight_data, int bits, int block_size, ncnn::Mat& weight_data_quantized, ncnn::Mat& weight_data_quantize_scales, ncnn::Mat& weight_data_dequantized) +static int quantize_weight(const ncnn::Mat& weight_data, int bits, int block_size, const ncnn::Mat& input_scales, ncnn::Mat& weight_data_quantized, ncnn::Mat& weight_data_quantize_scales, ncnn::Mat& weight_data_dequantized) { const int K = weight_data.w; const int N = weight_data.h; - const int packed_k_bytes = weight_quantize_packed_k_bytes(K, bits); const int block_count = (K + block_size - 1) / block_size; - weight_data_quantized.create(packed_k_bytes, N, (size_t)1u); + weight_data_quantized.create((K * bits + 7) / 8, N, (size_t)1u); weight_data_quantize_scales.create(block_count, N); weight_data_dequantized.create(K, N); if (weight_data_quantized.empty() || weight_data_quantize_scales.empty() || weight_data_dequantized.empty()) return -100; - memset(weight_data_quantized.data, 0, weight_data_quantized.total() * weight_data_quantized.elemsize); + weight_data_quantized.fill((unsigned char)0); const int qmax = (1 << (bits - 1)) - 1; + const float* input_scale_ptr = input_scales; for (int n = 0; n < N; n++) { const float* ptr = weight_data.row(n); - float* scale_ptr = weight_data_quantize_scales.row(n); unsigned char* qptr = weight_data_quantized.row(n); + float* scale_ptr = weight_data_quantize_scales.row(n); float* deqptr = weight_data_dequantized.row(n); for (int b = 0; b < block_count; b++) @@ -108,19 +114,19 @@ static int quantize_weight(const ncnn::Mat& weight_data, int bits, int block_siz float absmax = 0.f; for (int k = 0; k < max_kk; k++) { - const float v = (float)fabs(ptr[k0 + k]); + const float v = fabsf(ptr[k0 + k]); if (v > absmax) absmax = v; } - const float scale = absmax == 0.f ? 1.f : (float)qmax / absmax; + const float scale = absmax == 0.f ? 1.f : qmax / absmax; scale_ptr[b] = scale; for (int k = 0; k < max_kk; k++) { const int q = float2int_weight(ptr[k0 + k] * scale, bits); pack_signed_weight(qptr, k0 + k, bits, q); - deqptr[k0 + k] = q / scale; + deqptr[k0 + k] = q / scale * (input_scale_ptr ? input_scale_ptr[k0 + k] : 1.f); } } } @@ -128,85 +134,79 @@ static int quantize_weight(const ncnn::Mat& weight_data, int bits, int block_siz return 0; } -static int make_mha_weights(int qdim, int kdim, int vdim, int embed_dim, int bits, int block_size, std::vector& weights, std::vector& ref_weights, int input_scale = 0) +static int make_mha_weights(int qdim, int kdim, int vdim, int embed_dim, int bits, int block_size, int has_input_scale, std::vector& weights, std::vector& ref_weights) { - ncnn::Mat q_weight_data = RandomMat(qdim, embed_dim, -1.f, 1.f); - ncnn::Mat k_weight_data = RandomMat(kdim, embed_dim, -1.f, 1.f); - ncnn::Mat v_weight_data = RandomMat(vdim, embed_dim, -1.f, 1.f); - ncnn::Mat out_weight_data = RandomMat(embed_dim, qdim, -1.f, 1.f); - ncnn::Mat q_input_scales; ncnn::Mat k_input_scales; ncnn::Mat v_input_scales; ncnn::Mat out_input_scales; - - if (input_scale) + if (has_input_scale) { - q_input_scales = make_input_scales(qdim); - k_input_scales = make_input_scales(kdim); - v_input_scales = make_input_scales(vdim); - out_input_scales = make_input_scales(embed_dim); - - q_weight_data = scale_weight_by_input_scales(q_weight_data, q_input_scales, 1); - k_weight_data = scale_weight_by_input_scales(k_weight_data, k_input_scales, 1); - v_weight_data = scale_weight_by_input_scales(v_weight_data, v_input_scales, 1); - out_weight_data = scale_weight_by_input_scales(out_weight_data, out_input_scales, 1); + q_input_scales = make_input_scales(qdim, 0); + k_input_scales = make_input_scales(kdim, 1); + v_input_scales = make_input_scales(vdim, 2); + out_input_scales = make_input_scales(embed_dim, 3); + if (bits == 8) + out_input_scales.fill(1.f); } - ncnn::Mat q_weight_data_quantized; - ncnn::Mat k_weight_data_quantized; - ncnn::Mat v_weight_data_quantized; - ncnn::Mat out_weight_data_quantized; - ncnn::Mat q_weight_data_scales; - ncnn::Mat k_weight_data_scales; - ncnn::Mat v_weight_data_scales; - ncnn::Mat out_weight_data_scales; - ncnn::Mat q_weight_data_dequantized; - ncnn::Mat k_weight_data_dequantized; - ncnn::Mat v_weight_data_dequantized; - ncnn::Mat out_weight_data_dequantized; - - int ret = quantize_weight(q_weight_data, bits, block_size, q_weight_data_quantized, q_weight_data_scales, q_weight_data_dequantized); - if (ret != 0) - return ret; - ret = quantize_weight(k_weight_data, bits, block_size, k_weight_data_quantized, k_weight_data_scales, k_weight_data_dequantized); - if (ret != 0) - return ret; - ret = quantize_weight(v_weight_data, bits, block_size, v_weight_data_quantized, v_weight_data_scales, v_weight_data_dequantized); - if (ret != 0) - return ret; - ret = quantize_weight(out_weight_data, bits, block_size, out_weight_data_quantized, out_weight_data_scales, out_weight_data_dequantized); - if (ret != 0) - return ret; + ncnn::Mat q_weight = bits == 8 ? make_w8a8_mat(qdim, embed_dim, block_size) : RandomMat(qdim, embed_dim, -1.f, 1.f); + ncnn::Mat k_weight = bits == 8 ? make_w8a8_mat(kdim, embed_dim, block_size) : RandomMat(kdim, embed_dim, -1.f, 1.f); + ncnn::Mat v_weight = bits == 8 ? make_w8a8_mat(vdim, embed_dim, block_size) : RandomMat(vdim, embed_dim, -1.f, 1.f); + ncnn::Mat out_weight = bits == 8 ? make_w8a8_mat(embed_dim, qdim, block_size) : RandomMat(embed_dim, qdim, -1.f, 1.f); - if (input_scale) + if (bits == 8 && has_input_scale) { - q_weight_data_dequantized = scale_weight_by_input_scales(q_weight_data_dequantized, q_input_scales, 0); - k_weight_data_dequantized = scale_weight_by_input_scales(k_weight_data_dequantized, k_input_scales, 0); - v_weight_data_dequantized = scale_weight_by_input_scales(v_weight_data_dequantized, v_input_scales, 0); - out_weight_data_dequantized = scale_weight_by_input_scales(out_weight_data_dequantized, out_input_scales, 0); + const float* ptr = v_weight.row(0); + for (int i = 1; i < embed_dim; i++) + { + float* outptr = v_weight.row(i); + for (int j = 0; j < vdim; j++) + outptr[j] = ptr[j]; + } } - ncnn::Mat q_bias_data = RandomMat(embed_dim, -1.f, 1.f); - ncnn::Mat k_bias_data = RandomMat(embed_dim, -1.f, 1.f); - ncnn::Mat v_bias_data = RandomMat(embed_dim, -1.f, 1.f); - ncnn::Mat out_bias_data = RandomMat(qdim, -1.f, 1.f); - - weights.resize(input_scale ? 16 : 12); - weights[0] = q_weight_data_quantized; - weights[1] = q_bias_data; - weights[2] = k_weight_data_quantized; - weights[3] = k_bias_data; - weights[4] = v_weight_data_quantized; - weights[5] = v_bias_data; - weights[6] = out_weight_data_quantized; - weights[7] = out_bias_data; - weights[8] = q_weight_data_scales; - weights[9] = k_weight_data_scales; - weights[10] = v_weight_data_scales; - weights[11] = out_weight_data_scales; - - if (input_scale) + ncnn::Mat q_weight_quantized; + ncnn::Mat k_weight_quantized; + ncnn::Mat v_weight_quantized; + ncnn::Mat out_weight_quantized; + ncnn::Mat q_weight_scales; + ncnn::Mat k_weight_scales; + ncnn::Mat v_weight_scales; + ncnn::Mat out_weight_scales; + ncnn::Mat q_weight_dequantized; + ncnn::Mat k_weight_dequantized; + ncnn::Mat v_weight_dequantized; + ncnn::Mat out_weight_dequantized; + + if (quantize_weight(q_weight, bits, block_size, q_input_scales, q_weight_quantized, q_weight_scales, q_weight_dequantized) != 0 + || quantize_weight(k_weight, bits, block_size, k_input_scales, k_weight_quantized, k_weight_scales, k_weight_dequantized) != 0 + || quantize_weight(v_weight, bits, block_size, v_input_scales, v_weight_quantized, v_weight_scales, v_weight_dequantized) != 0 + || quantize_weight(out_weight, bits, block_size, out_input_scales, out_weight_quantized, out_weight_scales, out_weight_dequantized) != 0) + return -100; + + ncnn::Mat q_bias = RandomMat(embed_dim, -1.f, 1.f); + ncnn::Mat k_bias = RandomMat(embed_dim, -1.f, 1.f); + ncnn::Mat v_bias = RandomMat(embed_dim, -1.f, 1.f); + ncnn::Mat out_bias = RandomMat(qdim, -1.f, 1.f); + if (bits == 8 && has_input_scale) + v_bias.fill(0.f); + + weights.resize(has_input_scale ? 16 : 12); + weights[0] = q_weight_quantized; + weights[1] = q_bias; + weights[2] = k_weight_quantized; + weights[3] = k_bias; + weights[4] = v_weight_quantized; + weights[5] = v_bias; + weights[6] = out_weight_quantized; + weights[7] = out_bias; + weights[8] = q_weight_scales; + weights[9] = k_weight_scales; + weights[10] = v_weight_scales; + weights[11] = out_weight_scales; + + if (has_input_scale) { weights[12] = q_input_scales; weights[13] = k_input_scales; @@ -215,20 +215,84 @@ static int make_mha_weights(int qdim, int kdim, int vdim, int embed_dim, int bit } ref_weights.resize(8); - ref_weights[0] = q_weight_data_dequantized.reshape(embed_dim * qdim); - ref_weights[1] = q_bias_data; - ref_weights[2] = k_weight_data_dequantized.reshape(embed_dim * kdim); - ref_weights[3] = k_bias_data; - ref_weights[4] = v_weight_data_dequantized.reshape(embed_dim * vdim); - ref_weights[5] = v_bias_data; - ref_weights[6] = out_weight_data_dequantized.reshape(qdim * embed_dim); - ref_weights[7] = out_bias_data; + ref_weights[0] = q_weight_dequantized.reshape(embed_dim * qdim); + ref_weights[1] = q_bias; + ref_weights[2] = k_weight_dequantized.reshape(embed_dim * kdim); + ref_weights[3] = k_bias; + ref_weights[4] = v_weight_dequantized.reshape(embed_dim * vdim); + ref_weights[5] = v_bias; + ref_weights[6] = out_weight_dequantized.reshape(qdim * embed_dim); + ref_weights[7] = out_bias; return 0; } -static ncnn::ParamDict make_mha_param(int qdim, int kdim, int vdim, int embed_dim, int num_heads, int attn_mask, int kv_cache, int quantize_term) +static int test_multiheadattention_invalid_weight_block_quantize_term() { + const int invalid_quantize_terms[] = {403, 420, 700}; + + for (int i = 0; i < 3; i++) + { + ncnn::ParamDict pd; + pd.set(0, 8); + pd.set(1, 2); + pd.set(2, 64); + pd.set(3, 8); + pd.set(4, 8); + pd.set(18, invalid_quantize_terms[i]); + + ncnn::Layer* mha = ncnn::create_layer_naive("MultiHeadAttention"); + if (!mha) + return -100; + + const int ret = mha->load_param(pd); + delete mha; + + if (ret == 0) + { + fprintf(stderr, "test_multiheadattention_invalid_weight_block_quantize_term accepted quantize_term=%d\n", invalid_quantize_terms[i]); + return -1; + } + } + + return 0; +} + +static int test_multiheadattention_block_quant(int qdim, int kdim, int vdim, int embed_dim, int num_heads, int bits, int block_size, int attn_mask, int has_input_scale, int zero_input_group = 0) +{ + std::vector weights; + std::vector ref_weights; + int ret = make_mha_weights(qdim, kdim, vdim, embed_dim, bits, block_size, has_input_scale, weights, ref_weights); + if (ret != 0) + return ret; + + const int src_seqlen = 5; + const int dst_seqlen = 6; + std::vector as(3); + as[0] = bits == 8 ? make_w8a8_mat(qdim, src_seqlen, block_size, has_input_scale ? weights[12] : ncnn::Mat()) : RandomMat(qdim, src_seqlen, -1.f, 1.f); + as[1] = bits == 8 ? make_w8a8_mat(kdim, dst_seqlen, block_size, has_input_scale ? weights[13] : ncnn::Mat()) : RandomMat(kdim, dst_seqlen, -1.f, 1.f); + as[2] = bits == 8 ? make_w8a8_mat(vdim, dst_seqlen, block_size, has_input_scale ? weights[14] : ncnn::Mat()) : RandomMat(vdim, dst_seqlen, -1.f, 1.f); + + if (zero_input_group) + { + const int q_zero = qdim < block_size ? qdim : block_size; + const int k_zero = kdim < block_size ? kdim : block_size; + const int v_zero = vdim < block_size ? vdim : block_size; + for (int i = 0; i < src_seqlen; i++) + for (int j = 0; j < q_zero; j++) + as[0].row(i)[j] = 0.f; + for (int i = 0; i < dst_seqlen; i++) + { + for (int j = 0; j < k_zero; j++) + as[1].row(i)[j] = 0.f; + for (int j = 0; j < v_zero; j++) + as[2].row(i)[j] = 0.f; + } + } + + if (attn_mask) + as.push_back(RandomMat(dst_seqlen, src_seqlen, -1.f, 0.f)); + ncnn::ParamDict pd; pd.set(0, embed_dim); pd.set(1, num_heads); @@ -237,221 +301,380 @@ static ncnn::ParamDict make_mha_param(int qdim, int kdim, int vdim, int embed_di pd.set(4, vdim); pd.set(5, attn_mask); pd.set(6, 0.7f / sqrtf(embed_dim / num_heads)); - pd.set(7, kv_cache); - if (quantize_term) - pd.set(18, quantize_term); + pd.set(18, weight_block_quantize_term(bits, block_size, has_input_scale)); - return pd; -} - -static int run_mha_layer(const ncnn::ParamDict& pd, const std::vector& weights, const std::vector& inputs, int top_blob_count, std::vector& outputs) -{ ncnn::Option opt; - opt.num_threads = 2; opt.use_packing_layout = false; opt.use_fp16_packed = false; opt.use_fp16_storage = false; opt.use_fp16_arithmetic = false; opt.use_bf16_storage = false; - return test_layer_cpu(ncnn::LayerType::MultiHeadAttention, pd, weights, opt, inputs, top_blob_count, outputs, std::vector(), TEST_LAYER_DISABLE_GPU_TESTING); -} - -static int test_multiheadattention_block_quant(int qdim, int kdim, int vdim, int embed_dim, int num_heads, int bits, int block_size, int attn_mask, int input_scale = 0) -{ - std::vector weights; - std::vector ref_weights; - int ret = make_mha_weights(qdim, kdim, vdim, embed_dim, bits, block_size, weights, ref_weights, input_scale); - if (ret != 0) - return ret; - - std::vector inputs(3); - inputs[0] = RandomMat(qdim, 5, -1.f, 1.f); - inputs[1] = RandomMat(kdim, 6, -1.f, 1.f); - inputs[2] = RandomMat(vdim, 6, -1.f, 1.f); - - if (attn_mask) - inputs.push_back(RandomMat(6, 5, -1.f, 0.f)); + if (bits != 8) + { + ncnn::ParamDict ref_pd = pd; + ref_pd.set(18, 0); - const int quantize_term = weight_block_quantize_term(bits, block_size, input_scale); - const ncnn::ParamDict pd = make_mha_param(qdim, kdim, vdim, embed_dim, num_heads, attn_mask, 0, quantize_term); - const ncnn::ParamDict ref_pd = make_mha_param(qdim, kdim, vdim, embed_dim, num_heads, attn_mask, 0, 0); + std::vector refs; + ret = test_layer_naive(ncnn::layer_to_index("MultiHeadAttention"), ref_pd, ref_weights, as, 1, refs, TEST_LAYER_DISABLE_GPU_TESTING); - std::vector outputs; - std::vector refs; - ret = run_mha_layer(pd, weights, inputs, 1, outputs); - if (ret != 0) - { - fprintf(stderr, "test_multiheadattention_block_quant failed ret=%d qdim=%d kdim=%d vdim=%d embed_dim=%d bits=%d block_size=%d attn_mask=%d input_scale=%d\n", ret, qdim, kdim, vdim, embed_dim, bits, block_size, attn_mask, input_scale); - return ret; + for (int t = 0; t < 2 && ret == 0; t++) + { + std::vector outputs; + const int flags = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_cpu(ncnn::layer_to_index("MultiHeadAttention"), pd, weights, opt, as, 1, outputs, std::vector(), flags); + if (ret == 0) + ret = CompareMat(outputs, refs, 0.001f); + } } - - ret = run_mha_layer(ref_pd, ref_weights, inputs, 1, refs); - if (ret != 0) + else { - fprintf(stderr, "test_multiheadattention_block_quant reference failed ret=%d qdim=%d kdim=%d vdim=%d embed_dim=%d\n", ret, qdim, kdim, vdim, embed_dim); - return ret; + for (int t = 0; t < 2 && ret == 0; t++) + { + const int flags = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_opt("MultiHeadAttention", pd, weights, opt, as, 1, 0.001f, flags); + } } - - ret = CompareMat(outputs, refs, 0.001f); if (ret != 0) { - fprintf(stderr, "test_multiheadattention_block_quant compare failed qdim=%d kdim=%d vdim=%d embed_dim=%d bits=%d block_size=%d attn_mask=%d input_scale=%d\n", qdim, kdim, vdim, embed_dim, bits, block_size, attn_mask, input_scale); - return ret; + fprintf(stderr, "test_multiheadattention_block_quant failed qdim=%d kdim=%d vdim=%d embed_dim=%d heads=%d bits=%d block=%d mask=%d input_scale=%d zero=%d\n", qdim, kdim, vdim, embed_dim, num_heads, bits, block_size, attn_mask, has_input_scale, zero_input_group); } - return 0; + return ret; } -static int test_multiheadattention_block_quant_kvcache(int attn_mask = 0, int input_scale = 0) +static int test_multiheadattention_block_quant_kvcache(int bits, int block_size, int attn_mask, int has_input_scale) { const int qdim = 10; const int embed_dim = 8; - const int num_heads = 2; - const int bits = 4; - const int block_size = 64; + const int src_seqlen = bits == 8 ? 1 : 3; std::vector weights; std::vector ref_weights; - int ret = make_mha_weights(qdim, qdim, qdim, embed_dim, bits, block_size, weights, ref_weights, input_scale); + int ret = make_mha_weights(qdim, qdim, qdim, embed_dim, bits, block_size, has_input_scale, weights, ref_weights); if (ret != 0) return ret; - std::vector inputs(attn_mask ? 4 : 3); - inputs[0] = RandomMat(qdim, 3, -1.f, 1.f); + std::vector as(attn_mask ? 4 : 3); + as[0] = bits == 8 ? make_w8a8_mat(qdim, src_seqlen, block_size, has_input_scale ? weights[12] : ncnn::Mat()) : RandomMat(qdim, src_seqlen, -1.f, 1.f); if (attn_mask) { - inputs[1] = RandomMat(8, 3, -1.f, 0.f); - inputs[2] = RandomMat(5, embed_dim, -1.f, 1.f); - inputs[3] = RandomMat(5, embed_dim, -1.f, 1.f); + as[1] = RandomMat(5 + src_seqlen, src_seqlen, -1.f, 0.f); + as[2] = RandomMat(5, embed_dim, -1.f, 1.f); + as[3] = bits == 8 && has_input_scale ? make_w8a8_cache(5, embed_dim) : RandomMat(5, embed_dim, -1.f, 1.f); } else { - inputs[1] = RandomMat(5, embed_dim, -1.f, 1.f); - inputs[2] = RandomMat(5, embed_dim, -1.f, 1.f); + as[1] = RandomMat(5, embed_dim, -1.f, 1.f); + as[2] = bits == 8 && has_input_scale ? make_w8a8_cache(5, embed_dim) : RandomMat(5, embed_dim, -1.f, 1.f); } - const int quantize_term = weight_block_quantize_term(bits, block_size, input_scale); - const ncnn::ParamDict pd = make_mha_param(qdim, qdim, qdim, embed_dim, num_heads, attn_mask, 1, quantize_term); - const ncnn::ParamDict ref_pd = make_mha_param(qdim, qdim, qdim, embed_dim, num_heads, attn_mask, 1, 0); + ncnn::ParamDict pd; + pd.set(0, embed_dim); + pd.set(1, 2); + pd.set(2, embed_dim * qdim); + pd.set(3, qdim); + pd.set(4, qdim); + pd.set(5, attn_mask); + pd.set(6, 0.7f / sqrtf(4.f)); + pd.set(7, 1); + pd.set(18, weight_block_quantize_term(bits, block_size, has_input_scale)); - std::vector outputs; - std::vector refs; - ret = run_mha_layer(pd, weights, inputs, 3, outputs); - if (ret != 0) + ncnn::Option opt; + opt.use_packing_layout = false; + opt.use_fp16_packed = false; + opt.use_fp16_storage = false; + opt.use_fp16_arithmetic = false; + opt.use_bf16_storage = false; + + if (bits != 8) { - fprintf(stderr, "test_multiheadattention_block_quant_kvcache failed ret=%d attn_mask=%d input_scale=%d\n", ret, attn_mask, input_scale); - return ret; - } + ncnn::ParamDict ref_pd = pd; + ref_pd.set(18, 0); - ret = run_mha_layer(ref_pd, ref_weights, inputs, 3, refs); - if (ret != 0) + std::vector refs; + ret = test_layer_naive(ncnn::layer_to_index("MultiHeadAttention"), ref_pd, ref_weights, as, 3, refs, TEST_LAYER_DISABLE_GPU_TESTING); + + for (int t = 0; t < 2 && ret == 0; t++) + { + std::vector outputs; + const int flags = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_cpu(ncnn::layer_to_index("MultiHeadAttention"), pd, weights, opt, as, 3, outputs, std::vector(), flags); + if (ret == 0) + ret = CompareMat(outputs, refs, 0.001f); + } + } + else { - fprintf(stderr, "test_multiheadattention_block_quant_kvcache reference failed ret=%d attn_mask=%d input_scale=%d\n", ret, attn_mask, input_scale); - return ret; + for (int t = 0; t < 2 && ret == 0; t++) + { + const int flags = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_opt("MultiHeadAttention", pd, weights, opt, as, 3, 0.001f, flags); + } } - ret = CompareMat(outputs, refs, 0.001f); if (ret != 0) - { - fprintf(stderr, "test_multiheadattention_block_quant_kvcache compare failed attn_mask=%d input_scale=%d\n", attn_mask, input_scale); - return ret; - } + fprintf(stderr, "test_multiheadattention_block_quant_kvcache failed bits=%d block=%d mask=%d input_scale=%d\n", bits, block_size, attn_mask, has_input_scale); - return 0; + return ret; } -static int test_multiheadattention_block_quant_cross_kvcache(int attn_mask = 0, int input_scale = 0) +static int test_multiheadattention_block_quant_cross_kvcache(int bits, int block_size, int attn_mask, int has_input_scale) { const int qdim = 65; const int kdim = 33; const int vdim = 49; const int embed_dim = 64; - const int num_heads = 4; - const int bits = 6; - const int block_size = 32; + const int src_seqlen = bits == 8 ? 1 : 3; std::vector weights; std::vector ref_weights; - int ret = make_mha_weights(qdim, kdim, vdim, embed_dim, bits, block_size, weights, ref_weights, input_scale); + int ret = make_mha_weights(qdim, kdim, vdim, embed_dim, bits, block_size, has_input_scale, weights, ref_weights); if (ret != 0) return ret; - std::vector inputs(attn_mask ? 6 : 5); - inputs[0] = RandomMat(qdim, 3, -1.f, 1.f); - inputs[1] = RandomMat(kdim, 2, -1.f, 1.f); - inputs[2] = RandomMat(vdim, 2, -1.f, 1.f); + std::vector as(attn_mask ? 6 : 5); + as[0] = bits == 8 ? make_w8a8_mat(qdim, src_seqlen, block_size, has_input_scale ? weights[12] : ncnn::Mat()) : RandomMat(qdim, src_seqlen, -1.f, 1.f); + as[1] = bits == 8 ? make_w8a8_mat(kdim, 2, block_size, has_input_scale ? weights[13] : ncnn::Mat()) : RandomMat(kdim, 2, -1.f, 1.f); + as[2] = bits == 8 ? make_w8a8_mat(vdim, 2, block_size, has_input_scale ? weights[14] : ncnn::Mat()) : RandomMat(vdim, 2, -1.f, 1.f); if (attn_mask) { - inputs[3] = RandomMat(5, 3, -1.f, 0.f); - inputs[4] = RandomMat(5, embed_dim, -1.f, 1.f); - inputs[5] = RandomMat(5, embed_dim, -1.f, 1.f); + as[3] = RandomMat(5, src_seqlen, -1.f, 0.f); + as[4] = RandomMat(5, embed_dim, -1.f, 1.f); + as[5] = bits == 8 && has_input_scale ? make_w8a8_cache(5, embed_dim) : RandomMat(5, embed_dim, -1.f, 1.f); } else { - inputs[3] = RandomMat(5, embed_dim, -1.f, 1.f); - inputs[4] = RandomMat(5, embed_dim, -1.f, 1.f); + as[3] = RandomMat(5, embed_dim, -1.f, 1.f); + as[4] = bits == 8 && has_input_scale ? make_w8a8_cache(5, embed_dim) : RandomMat(5, embed_dim, -1.f, 1.f); } - const int quantize_term = weight_block_quantize_term(bits, block_size, input_scale); - const ncnn::ParamDict pd = make_mha_param(qdim, kdim, vdim, embed_dim, num_heads, attn_mask, 1, quantize_term); - const ncnn::ParamDict ref_pd = make_mha_param(qdim, kdim, vdim, embed_dim, num_heads, attn_mask, 1, 0); + ncnn::ParamDict pd; + pd.set(0, embed_dim); + pd.set(1, 4); + pd.set(2, embed_dim * qdim); + pd.set(3, kdim); + pd.set(4, vdim); + pd.set(5, attn_mask); + pd.set(6, 0.7f / sqrtf(16.f)); + pd.set(7, 1); + pd.set(18, weight_block_quantize_term(bits, block_size, has_input_scale)); - std::vector outputs; - std::vector refs; - ret = run_mha_layer(pd, weights, inputs, 3, outputs); - if (ret != 0) + ncnn::Option opt; + opt.use_packing_layout = false; + opt.use_fp16_packed = false; + opt.use_fp16_storage = false; + opt.use_fp16_arithmetic = false; + opt.use_bf16_storage = false; + + if (bits != 8) { - fprintf(stderr, "test_multiheadattention_block_quant_cross_kvcache failed ret=%d attn_mask=%d input_scale=%d\n", ret, attn_mask, input_scale); - return ret; + ncnn::ParamDict ref_pd = pd; + ref_pd.set(18, 0); + + std::vector refs; + ret = test_layer_naive(ncnn::layer_to_index("MultiHeadAttention"), ref_pd, ref_weights, as, 3, refs, TEST_LAYER_DISABLE_GPU_TESTING); + + for (int t = 0; t < 2 && ret == 0; t++) + { + std::vector outputs; + const int flags = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_cpu(ncnn::layer_to_index("MultiHeadAttention"), pd, weights, opt, as, 3, outputs, std::vector(), flags); + if (ret == 0) + ret = CompareMat(outputs, refs, 0.001f); + } + } + else + { + for (int t = 0; t < 2 && ret == 0; t++) + { + const int flags = TEST_LAYER_DISABLE_GPU_TESTING | (t ? TEST_LAYER_ENABLE_THREADING : 0); + ret = test_layer_opt("MultiHeadAttention", pd, weights, opt, as, 3, 0.001f, flags); + } } - ret = run_mha_layer(ref_pd, ref_weights, inputs, 3, refs); if (ret != 0) - { - fprintf(stderr, "test_multiheadattention_block_quant_cross_kvcache reference failed ret=%d attn_mask=%d input_scale=%d\n", ret, attn_mask, input_scale); + fprintf(stderr, "test_multiheadattention_block_quant_cross_kvcache failed bits=%d block=%d mask=%d input_scale=%d\n", bits, block_size, attn_mask, has_input_scale); + + return ret; +} + +static int test_multiheadattention_block_quant_pipeline() +{ + const int qdim = 35; + const int embed_dim = 32; + const int block_size = 32; + + std::vector weights; + std::vector ref_weights; + int ret = make_mha_weights(qdim, qdim, qdim, embed_dim, 8, block_size, 1, weights, ref_weights); + if (ret != 0) return ret; - } - ret = CompareMat(outputs, refs, 0.001f); + ncnn::ParamDict pd; + pd.set(0, embed_dim); + pd.set(1, 4); + pd.set(2, embed_dim * qdim); + pd.set(3, qdim); + pd.set(4, qdim); + pd.set(6, 0.7f / sqrtf(8.f)); + pd.set(7, 1); + pd.set(18, weight_block_quantize_term(8, block_size, 1)); + + ncnn::Option opt; + opt.lightmode = false; + opt.num_threads = 2; + opt.use_packing_layout = false; + opt.use_fp16_packed = false; + opt.use_fp16_storage = false; + opt.use_fp16_arithmetic = false; + opt.use_bf16_storage = false; + + ncnn::Layer* mha = ncnn::create_layer_cpu("MultiHeadAttention"); + if (!mha) + return -100; + + ret = mha->load_param(pd); + if (ret == 0) + ret = mha->load_model(ncnn::ModelBinFromMatArray(weights.data())); + if (ret == 0) + ret = mha->create_pipeline(opt); if (ret != 0) { - fprintf(stderr, "test_multiheadattention_block_quant_cross_kvcache compare failed attn_mask=%d input_scale=%d\n", attn_mask, input_scale); + delete mha; return ret; } - return 0; + int test_ret = 0; + std::vector prefill_inputs(3); + prefill_inputs[0] = make_w8a8_mat(qdim, 4, block_size, weights[12]); + + std::vector prefill_reference; + std::vector prefill_outputs(3); + test_ret = test_layer_naive(ncnn::layer_to_index("MultiHeadAttention"), pd, weights, prefill_inputs, 3, prefill_reference, TEST_LAYER_DISABLE_GPU_TESTING); + if (test_ret == 0) + test_ret = mha->forward(prefill_inputs, prefill_outputs, opt); + if (test_ret == 0) + test_ret = CompareMat(prefill_outputs, prefill_reference, 0.001f); + + for (int i = 0; test_ret == 0 && i < 3; i++) + { + if (prefill_outputs[i].elembits() != 32 || prefill_outputs[i].elempack != 1) + test_ret = -1; + } + if (test_ret == 0 && (prefill_outputs[1].w != 4 || prefill_outputs[2].w != 4)) + test_ret = -1; + + std::vector decode_reference_inputs(3); + std::vector decode_inputs(3); + decode_inputs[0] = make_w8a8_mat(qdim, 1, block_size, weights[12]); + decode_reference_inputs[0] = decode_inputs[0]; + if (test_ret == 0) + { + decode_reference_inputs[1] = prefill_reference[1]; + decode_reference_inputs[2] = prefill_reference[2]; + decode_inputs[1] = prefill_outputs[1]; + decode_inputs[2] = prefill_outputs[2]; + } + + std::vector decode_reference; + std::vector decode_outputs(3); + if (test_ret == 0) + test_ret = test_layer_naive(ncnn::layer_to_index("MultiHeadAttention"), pd, weights, decode_reference_inputs, 3, decode_reference, TEST_LAYER_DISABLE_GPU_TESTING); + ncnn::Option decode_opt = opt; + decode_opt.num_threads = 4; + if (test_ret == 0) + test_ret = mha->forward(decode_inputs, decode_outputs, decode_opt); + if (test_ret == 0) + test_ret = CompareMat(decode_outputs, decode_reference, 0.001f); + + for (int i = 0; test_ret == 0 && i < 3; i++) + { + if (decode_outputs[i].elembits() != 32 || decode_outputs[i].elempack != 1) + test_ret = -1; + } + if (test_ret == 0 && (decode_outputs[1].w != 5 || decode_outputs[2].w != 5)) + test_ret = -1; + + const int destroy_ret = mha->destroy_pipeline(opt); + delete mha; + + if (test_ret != 0) + { + fprintf(stderr, "test_multiheadattention_block_quant_pipeline failed ret=%d\n", test_ret); + return test_ret; + } + + return destroy_ret; } +static int test_multiheadattention_block_quant_0() +{ + return 0 + || test_multiheadattention_block_quant(13, 9, 11, 8, 2, 4, 32, 0, 0) + || test_multiheadattention_block_quant(10, 10, 10, 8, 2, 6, 64, 1, 0) + || test_multiheadattention_block_quant(12, 7, 9, 8, 2, 8, 128, 0, 0) + || test_multiheadattention_block_quant(35, 33, 31, 32, 4, 8, 32, 1, 0) + || test_multiheadattention_block_quant(35, 33, 31, 72, 8, 8, 32, 0, 0) + || test_multiheadattention_block_quant(65, 33, 49, 64, 4, 8, 64, 0, 1) + || test_multiheadattention_block_quant(129, 129, 129, 128, 8, 8, 128, 1, 1) + || test_multiheadattention_block_quant(13, 9, 11, 8, 2, 4, 64, 1, 1) + || test_multiheadattention_block_quant(35, 33, 31, 32, 4, 8, 32, 0, 0, 1); +} + +static int test_multiheadattention_block_quant_1() +{ + return 0 + || test_multiheadattention_block_quant_kvcache(4, 64, 0, 0) + || test_multiheadattention_block_quant_kvcache(4, 32, 1, 1) + || test_multiheadattention_block_quant_kvcache(8, 32, 0, 0) + || test_multiheadattention_block_quant_kvcache(8, 64, 1, 1); +} + +static int test_multiheadattention_block_quant_2() +{ + return 0 + || test_multiheadattention_block_quant_cross_kvcache(6, 32, 0, 0) + || test_multiheadattention_block_quant_cross_kvcache(6, 64, 1, 1) + || test_multiheadattention_block_quant_cross_kvcache(8, 32, 0, 0) + || test_multiheadattention_block_quant_cross_kvcache(8, 128, 1, 1); +} + +#endif // NCNN_WEIGHT_QUANT + int main() { SRAND(7767517); -#if !NCNN_WEIGHT_QUANT - ncnn::ParamDict pd = make_mha_param(5, 5, 5, 4, 2, 0, 0, 410); +#if NCNN_WEIGHT_QUANT + return 0 + || test_multiheadattention_invalid_weight_block_quantize_term() + || test_multiheadattention_block_quant_0() + || test_multiheadattention_block_quant_1() + || test_multiheadattention_block_quant_2() + || test_multiheadattention_block_quant_pipeline(); +#else + ncnn::ParamDict pd; + pd.set(0, 4); + pd.set(1, 2); + pd.set(2, 20); + pd.set(3, 5); + pd.set(4, 5); + pd.set(18, 410); + + ncnn::Layer* mha = ncnn::create_layer_naive("MultiHeadAttention"); + if (!mha) + return -100; + + const int ret = mha->load_param(pd); + delete mha; - ncnn::MultiHeadAttention mha; - if (mha.load_param(pd) == 0) + if (ret == 0) { fprintf(stderr, "test_multiheadattention_block_quant failed NCNN_WEIGHT_QUANT=OFF accepted weight block quantization\n"); return -1; } return 0; -#else - return 0 - || test_multiheadattention_block_quant(13, 9, 11, 8, 2, 4, 32, 0) - || test_multiheadattention_block_quant(10, 10, 10, 8, 2, 6, 64, 1) - || test_multiheadattention_block_quant(12, 7, 9, 8, 2, 8, 128, 0) - || test_multiheadattention_block_quant(13, 9, 11, 8, 2, 4, 64, 1, 1) - || test_multiheadattention_block_quant(65, 65, 65, 64, 4, 6, 64, 1) - || test_multiheadattention_block_quant(65, 33, 49, 64, 4, 4, 32, 0, 1) - || test_multiheadattention_block_quant_kvcache() - || test_multiheadattention_block_quant_kvcache(0, 1) - || test_multiheadattention_block_quant_kvcache(1) - || test_multiheadattention_block_quant_kvcache(1, 1) - || test_multiheadattention_block_quant_cross_kvcache() - || test_multiheadattention_block_quant_cross_kvcache(1) - || test_multiheadattention_block_quant_cross_kvcache(0, 1); #endif } diff --git a/tests/test_multiheadattention_oom.cpp b/tests/test_multiheadattention_oom.cpp index 86ff50e1c5d..6851361ea97 100644 --- a/tests/test_multiheadattention_oom.cpp +++ b/tests/test_multiheadattention_oom.cpp @@ -55,9 +55,66 @@ static int test_multiheadattention_0() || test_multiheadattention_oom(RandomMat(12, 17), RandomMat(28, 32), RandomMat(11, 32), 12, 3, 1); } +#if NCNN_WEIGHT_QUANT +static int test_multiheadattention_w8a8_oom(int qdim, int kdim, int vdim, int embed_dim, int num_heads, int block_size, int input_scale) +{ + const int block_size_code = block_size == 32 ? 0 : block_size == 64 ? 1 : 2; + + ncnn::ParamDict pd; + pd.set(0, embed_dim); + pd.set(1, num_heads); + pd.set(2, embed_dim * qdim); + pd.set(3, kdim); + pd.set(4, vdim); + pd.set(5, 0); + pd.set(18, 800 + input_scale * 10 + block_size_code); + + std::vector weights(input_scale ? 16 : 12); + weights[0] = RandomS8Mat(qdim, embed_dim); + weights[1] = RandomMat(embed_dim); + weights[2] = RandomS8Mat(kdim, embed_dim); + weights[3] = RandomMat(embed_dim); + weights[4] = RandomS8Mat(vdim, embed_dim); + weights[5] = RandomMat(embed_dim); + weights[6] = RandomS8Mat(embed_dim, qdim); + weights[7] = RandomMat(qdim); + weights[8] = RandomMat((qdim + block_size - 1) / block_size, embed_dim, 10.f, 20.f); + weights[9] = RandomMat((kdim + block_size - 1) / block_size, embed_dim, 10.f, 20.f); + weights[10] = RandomMat((vdim + block_size - 1) / block_size, embed_dim, 10.f, 20.f); + weights[11] = RandomMat((embed_dim + block_size - 1) / block_size, qdim, 10.f, 20.f); + + if (input_scale) + { + weights[12] = RandomMat(qdim, 0.5f, 1.5f); + weights[13] = RandomMat(kdim, 0.5f, 1.5f); + weights[14] = RandomMat(vdim, 0.5f, 1.5f); + weights[15] = RandomMat(embed_dim, 0.5f, 1.5f); + } + + std::vector inputs(3); + inputs[0] = RandomMat(qdim, 3); + inputs[1] = RandomMat(kdim, 5); + inputs[2] = RandomMat(vdim, 5); + + int ret = test_layer_oom("MultiHeadAttention", pd, weights, inputs, 1, TEST_LAYER_ENABLE_THREADING); + if (ret != 0) + { + fprintf(stderr, "test_multiheadattention_w8a8_oom failed qdim=%d kdim=%d vdim=%d embed_dim=%d num_heads=%d block_size=%d input_scale=%d\n", qdim, kdim, vdim, embed_dim, num_heads, block_size, input_scale); + } + + return ret; +} +#endif // NCNN_WEIGHT_QUANT + int main() { SRAND(7767517); - return test_multiheadattention_0(); + return 0 + || test_multiheadattention_0() +#if NCNN_WEIGHT_QUANT + || test_multiheadattention_w8a8_oom(33, 35, 37, 8, 2, 32, 0) + || test_multiheadattention_w8a8_oom(65, 67, 69, 8, 2, 64, 1) +#endif + ; }