diff --git a/docs/2026.html b/docs/2026.html
index 72069d4ec4..57cf7ebc1b 100644
--- a/docs/2026.html
+++ b/docs/2026.html
@@ -51,6 +51,7 @@
New features
SVE2 optimizations of class SynetConvolution32fNhwcDirect.
SVE2 optimizations of class SynetInnerProduct32fGemm.
SVE2 optimizations of class SynetInnerProduct32fProd.
+ SVE2 optimizations of class SynetInnerProduct16bGemmNN.
Renaming
diff --git a/prj/vs2022/Sve2.vcxproj b/prj/vs2022/Sve2.vcxproj
index 07680b795c..3a80c5e7d1 100644
--- a/prj/vs2022/Sve2.vcxproj
+++ b/prj/vs2022/Sve2.vcxproj
@@ -128,6 +128,8 @@
+
+
diff --git a/prj/vs2022/Sve2.vcxproj.filters b/prj/vs2022/Sve2.vcxproj.filters
index cb208ec8d2..d8137e9576 100644
--- a/prj/vs2022/Sve2.vcxproj.filters
+++ b/prj/vs2022/Sve2.vcxproj.filters
@@ -545,6 +545,12 @@
Sve2\Synet\MergedConvolution
+
+ Sve2\Synet\InnerProduct
+
+
+ Sve2\Synet\InnerProduct
+
Sve2\Synet\InnerProduct
diff --git a/src/Simd/SimdLib.cpp b/src/Simd/SimdLib.cpp
index 9c1fee0617..41f4ed122c 100644
--- a/src/Simd/SimdLib.cpp
+++ b/src/Simd/SimdLib.cpp
@@ -6773,7 +6773,7 @@ SIMD_API void* SimdSynetInnerProduct16bInit(size_t M, size_t N, size_t K, SimdTe
SIMD_EMPTY();
#if defined(SIMD_SYNET_ENABLE)
typedef void* (*SimdSynetInnerProduct16bInitPtr) (size_t M, size_t N, size_t K, SimdTensorDataType typeA, SimdTensorDataType typeB, SimdTensorDataType typeC, SimdBool transB, SimdBool constB, SimdBool bias, SimdConvolutionActivationType activation);
- const static SimdSynetInnerProduct16bInitPtr simdSynetInnerProduct16bInit = SIMD_FUNC4(SynetInnerProduct16bInit, SIMD_AMXBF16_FUNC, SIMD_AVX512BW_FUNC, SIMD_AVX2_FUNC, SIMD_SSE41_FUNC);
+ const static SimdSynetInnerProduct16bInitPtr simdSynetInnerProduct16bInit = SIMD_FUNC5(SynetInnerProduct16bInit, SIMD_AMXBF16_FUNC, SIMD_AVX512BW_FUNC, SIMD_AVX2_FUNC, SIMD_SSE41_FUNC, SIMD_SVE2_FUNC);
return simdSynetInnerProduct16bInit(M, N, K, typeA, typeB, typeC, transB, constB, bias, activation);
#else
diff --git a/src/Simd/SimdSve2SynetInnerProduct16b.cpp b/src/Simd/SimdSve2SynetInnerProduct16b.cpp
new file mode 100644
index 0000000000..3264f3488a
--- /dev/null
+++ b/src/Simd/SimdSve2SynetInnerProduct16b.cpp
@@ -0,0 +1,42 @@
+/*
+* Simd Library (http://ermig1979.github.io/Simd).
+*
+* Copyright (c) 2011-2026 Yermalayeu Ihar.
+*
+* Permission is hereby granted, free of charge, to any person obtaining a copy
+* of this software and associated documentation files (the "Software"), to deal
+* in the Software without restriction, including without limitation the rights
+* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+* copies of the Software, and to permit persons to whom the Software is
+* furnished to do so, subject to the following conditions:
+*
+* The above copyright notice and this permission notice shall be included in
+* all copies or substantial portions of the Software.
+*
+* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+* SOFTWARE.
+*/
+#include "Simd/SimdSynetInnerProduct16b.h"
+
+namespace Simd
+{
+#if defined(SIMD_SVE2_ENABLE) && defined(SIMD_SYNET_ENABLE)
+ namespace Sve2
+ {
+ void* SynetInnerProduct16bInit(size_t M, size_t N, size_t K, SimdTensorDataType typeA, SimdTensorDataType typeB, SimdTensorDataType typeC, SimdBool transB, SimdBool constB, SimdBool bias, SimdConvolutionActivationType activation)
+ {
+ InnerProductParam16b param(M, N, K, typeA, typeB, typeC, transB, constB, bias, activation);
+ if (!param.Valid())
+ return NULL;
+ if (Base::SynetInnerProduct16bGemmNN::Preferable(param))
+ return new Sve2::SynetInnerProduct16bGemmNN(param);
+ return Base::SynetInnerProduct16bInit(M, N, K, typeA, typeB, typeC, transB, constB, bias, activation);
+ }
+ }
+#endif
+}
diff --git a/src/Simd/SimdSve2SynetInnerProduct16bGemmNN.cpp b/src/Simd/SimdSve2SynetInnerProduct16bGemmNN.cpp
new file mode 100644
index 0000000000..d3153d527a
--- /dev/null
+++ b/src/Simd/SimdSve2SynetInnerProduct16bGemmNN.cpp
@@ -0,0 +1,644 @@
+/*
+* Simd Library (http://ermig1979.github.io/Simd).
+*
+* Copyright (c) 2011-2026 Yermalayeu Ihar.
+*
+* Permission is hereby granted, free of charge, to any person obtaining a copy
+* of this software and associated documentation files (the "Software"), to deal
+* in the Software without restriction, including without limitation the rights
+* to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+* copies of the Software, and to permit persons to whom the Software is
+* furnished to do so, subject to the following conditions:
+*
+* The above copyright notice and this permission notice shall be included in
+* all copies or substantial portions of the Software.
+*
+* THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+* IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+* FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+* AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+* LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+* OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+* SOFTWARE.
+*/
+#include "Simd/SimdMemory.h"
+#include "Simd/SimdStore.h"
+#include "Simd/SimdSynetInnerProduct16b.h"
+#include "Simd/SimdSynetConvolution16bCommon.h"
+#include "Simd/SimdSynetActivation.h"
+#include "Simd/SimdSynet.h"
+#include "Simd/SimdBase.h"
+#include "Simd/SimdSve2.h"
+#include "Simd/SimdCpu.h"
+#include "Simd/SimdBFloat16.h"
+
+namespace Simd
+{
+#if defined(SIMD_SVE2_ENABLE) && defined(SIMD_SYNET_ENABLE)
+ namespace Sve2
+ {
+ typedef Base::SynetInnerProduct16bGemmNN::AlgParam AlgParam;
+ typedef Base::SynetInnerProduct16bGemmNN::GemmPtr GemmPtr;
+
+ //-----------------------------------------------------------------------------------------
+
+ SIMD_INLINE svuint32_t Float32ToBFloat16(svfloat32_t value, const svbool_t& mask)
+ {
+ svuint32_t bits = svreinterpret_u32_f32(value);
+ svuint32_t round = svadd_n_u32_x(mask, svand_n_u32_x(mask, svlsr_n_u32_x(mask, bits, Base::Bf16::SHIFT), 1), Base::Bf16::ROUND);
+ return svlsr_n_u32_x(mask, svadd_u32_x(mask, bits, round), Base::Bf16::SHIFT);
+ }
+
+ SIMD_INLINE svfloat32_t BroadcastBf16(uint16_t value)
+ {
+ return svreinterpret_f32_u32(svdup_n_u32(uint32_t(value) << Base::Bf16::SHIFT));
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ static void InnerProduct16bGemmNN_ConvertA(const uint8_t* src8, const InnerProductParam16b& p, const AlgParam& a, size_t M, size_t K, uint16_t* dst)
+ {
+ const float* src = (float*)src8;
+ if (p.K == a.aK)
+ {
+ Float32ToBFloat16(src, K * M, dst);
+ }
+ else
+ {
+ for (size_t i = 0; i < M; ++i)
+ {
+ Float32ToBFloat16(src, p.K, dst);
+ for (size_t k = p.K; k < a.aK; ++k)
+ dst[k] = 0;
+ src += p.K;
+ dst += a.aK;
+ }
+ }
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ static void InnerProduct16bGemmNN_ReorderA(const uint8_t* src8, const InnerProductParam16b& p, const AlgParam& a, size_t M, size_t K, uint16_t* dst)
+ {
+ const uint16_t* src = (uint16_t*)src8;
+ for (size_t i = 0; i < M; ++i)
+ {
+ memcpy(dst, src, p.K * sizeof(uint16_t));
+ for (size_t k = p.K; k < a.aK; ++k)
+ dst[k] = 0;
+ src += p.K;
+ dst += a.aK;
+ }
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ SIMD_INLINE void ConvertBn(const float* src, size_t stride, uint16_t* dst, const svbool_t& mask)
+ {
+ svfloat32_t s0 = svld1_f32(mask, src + 0 * stride);
+ svfloat32_t s1 = svld1_f32(mask, src + 1 * stride);
+ svuint32_t d0 = Float32ToBFloat16(s0, mask);
+ svuint32_t d1 = svlsl_n_u32_x(mask, Float32ToBFloat16(s1, mask), Base::Bf16::SHIFT);
+ svst1_u32(mask, (uint32_t*)dst, svorr_u32_x(mask, d0, d1));
+ }
+
+ static void InnerProduct16bGemmNN_ConvertBn(const uint8_t* src8, const InnerProductParam16b& p, const AlgParam& a, size_t N, size_t K, uint16_t* dst)
+ {
+ const float* src = (float*)src8;
+ const size_t F = a.F;
+ const svbool_t body = svptrue_b32();
+ size_t Kl = AlignLo(K, a.microK), Kh = AlignHi(K, a.microK), Nf = AlignLo(N, a.F), j = 0, gap = (a.bK - Kh) * a.F;
+ for (; j < Nf; j += a.F)
+ {
+ size_t k = 0;
+ for (; k < Kl; k += 2)
+ {
+ const float* ps = src + k * p.N + j;
+ ConvertBn(ps, p.N, dst, body);
+ dst += F * 2;
+ }
+ for (; k < Kh; k += 2)
+ {
+ const float* ps = src + k * p.N + j;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = Base::Float32ToBFloat16(ps[i * p.N + f]);
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ dst += gap;
+ }
+ for (; j < N; j += a.F)
+ {
+ for (size_t k = 0; k < Kh; k += 2)
+ {
+ const float* ps = src + k * p.N + j;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = Base::Float32ToBFloat16(ps[i * p.N + f]);
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ }
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ SIMD_INLINE void ConvertBt(const float* src, size_t stride, uint16_t* dst, size_t F)
+ {
+ for (size_t f = 0; f < F; ++f)
+ {
+ dst[0] = Base::Float32ToBFloat16(src[f * stride + 0]);
+ dst[1] = Base::Float32ToBFloat16(src[f * stride + 1]);
+ dst += 2;
+ }
+ }
+
+ static void InnerProduct16bGemmNN_ConvertBt(const uint8_t* src8, const InnerProductParam16b& p, const AlgParam& a, size_t N, size_t K, uint16_t* dst)
+ {
+ const float* src = (float*)src8;
+ const size_t F = a.F;
+ size_t Kl = AlignLo(K, a.microK), Kh = AlignHi(K, a.microK), Nf = AlignLo(N, a.F), j = 0, gap = (a.bK - Kh) * a.F;
+ for (; j < Nf; j += a.F)
+ {
+ size_t k = 0;
+ for (; k < Kl; k += 2)
+ {
+ const float* ps = src + j * p.K + k;
+ ConvertBt(ps, p.K, dst, F);
+ dst += F * 2;
+ }
+ for (; k < Kh; k += 2)
+ {
+ const float* ps = src + j * p.K + k;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = Base::Float32ToBFloat16(ps[f * p.K + i]);
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ dst += gap;
+ }
+ for (; j < N; j += a.F)
+ {
+ for (size_t k = 0; k < Kh; k += 2)
+ {
+ const float* ps = src + j * p.K + k;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = Base::Float32ToBFloat16(ps[f * p.K + i]);
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ }
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ SIMD_INLINE void ReorderBn(const uint16_t* src, size_t stride, uint16_t* dst, const svbool_t& mask)
+ {
+ svuint32_t d0 = svld1uh_u32(mask, src + 0 * stride);
+ svuint32_t d1 = svld1uh_u32(mask, src + 1 * stride);
+ svst1_u32(mask, (uint32_t*)dst, svorr_u32_x(mask, d0, svlsl_n_u32_x(mask, d1, Base::Bf16::SHIFT)));
+ }
+
+ static void InnerProduct16bGemmNN_ReorderBn(const uint8_t* src8, const InnerProductParam16b& p, const AlgParam& a, size_t N, size_t K, uint16_t* dst)
+ {
+ const uint16_t* src = (uint16_t*)src8;
+ const size_t F = a.F;
+ const svbool_t body = svptrue_b32();
+ size_t Kl = AlignLo(K, a.microK), Kh = AlignHi(K, a.microK), Nf = AlignLo(N, a.F), j = 0, gap = (a.bK - Kh) * a.F;
+ for (; j < Nf; j += a.F)
+ {
+ size_t k = 0;
+ for (; k < Kl; k += 2)
+ {
+ const uint16_t* ps = src + k * p.N + j;
+ ReorderBn(ps, p.N, dst, body);
+ dst += F * 2;
+ }
+ for (; k < Kh; k += 2)
+ {
+ const uint16_t* ps = src + k * p.N + j;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = ps[i * p.N + f];
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ dst += gap;
+ }
+ for (; j < N; j += a.F)
+ {
+ for (size_t k = 0; k < Kh; k += 2)
+ {
+ const uint16_t* ps = src + k * p.N + j;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = ps[i * p.N + f];
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ }
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ SIMD_INLINE void ReorderBt(const uint16_t* src, size_t stride, uint16_t* dst, size_t F)
+ {
+ for (size_t f = 0; f < F; ++f)
+ {
+ ((uint32_t*)dst)[f] = ((uint32_t*)(src + f * stride))[0];
+ }
+ }
+
+ static void InnerProduct16bGemmNN_ReorderBt(const uint8_t* src8, const InnerProductParam16b& p, const AlgParam& a, size_t N, size_t K, uint16_t* dst)
+ {
+ const uint16_t* src = (uint16_t*)src8;
+ const size_t F = a.F;
+ size_t Kl = AlignLo(K, a.microK), Kh = AlignHi(K, a.microK), Nf = AlignLo(N, a.F), j = 0, gap = (a.bK - Kh) * a.F;
+ for (; j < Nf; j += a.F)
+ {
+ size_t k = 0;
+ for (; k < Kl; k += 2)
+ {
+ const uint16_t* ps = src + j * p.K + k;
+ ReorderBt(ps, p.K, dst, F);
+ dst += F * 2;
+ }
+ for (; k < Kh; k += 2)
+ {
+ const uint16_t* ps = src + j * p.K + k;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = ps[f * p.K + i];
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ dst += gap;
+ }
+ for (; j < N; j += a.F)
+ {
+ for (size_t k = 0; k < Kh; k += 2)
+ {
+ const uint16_t* ps = src + j * p.K + k;
+ for (size_t f = 0; f < a.F; ++f)
+ {
+ for (size_t i = 0; i < 2; ++i)
+ {
+ if (j + f < N && k + i < K)
+ *(dst++) = ps[f * p.K + i];
+ else
+ *(dst++) = 0;
+ }
+ }
+ }
+ }
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ template SIMD_INLINE void Save1(
+ uint8_t* dst, float* buf, svfloat32_t val, svfloat32_t bias, svfloat32_t param0, svfloat32_t param1, size_t index, const svbool_t& mask)
+ {
+ if (term == Term16bInterim)
+ {
+ svst1_f32(mask, buf, val);
+ }
+ else
+ {
+ svfloat32_t f32 = Activate(svadd_f32_x(mask, val, bias), param0, param1, index, mask);
+ if (term == Term16bLast16b)
+ svst1h_u32(mask, (uint16_t*)dst, Float32ToBFloat16(f32, mask));
+ else
+ svst1_f32(mask, (float*)dst, f32);
+ }
+ }
+
+ template SIMD_INLINE void Save2(
+ uint8_t* dst, float* buf, svfloat32_t val0, svfloat32_t val1,
+ svfloat32_t bias0, svfloat32_t bias1, svfloat32_t param0, svfloat32_t param1,
+ const svbool_t& mask0, const svbool_t& mask1, size_t F)
+ {
+ Save1(dst, buf, val0, bias0, param0, param1, 0, mask0);
+ if (term == Term16bInterim)
+ Save1(dst, buf + F, val1, bias1, param0, param1, 1, mask1);
+ else if (term == Term16bLast16b)
+ Save1(dst + F * sizeof(uint16_t), buf, val1, bias1, param0, param1, 1, mask1);
+ else
+ Save1(dst + F * sizeof(float), buf, val1, bias1, param0, param1, 1, mask1);
+ }
+
+ //-----------------------------------------------------------------------------------------
+
+ template void InnerProduct16bGemmNN_2xM(
+ const uint16_t* A0, const InnerProductParam16b& p, const AlgParam& a,
+ size_t N, size_t K, int update, const uint16_t* B0, float* C,
+ svfloat32_t bias0, svfloat32_t bias1, svfloat32_t param0, svfloat32_t param1, uint8_t* dst)
+ {
+ const size_t F = a.F, DF = F * 2;
+ const svbool_t body = svptrue_b32();
+ svfloat32_t c00, c01, c10, c11, c20, c21, c30, c31, c40, c41, a0, b00, b01, b10, b11;
+ svuint32_t maskBits = svdup_n_u32(Base::Bf16::MASK);
+ size_t dC = a.cN, dA = a.aK, dD = p.N * a.eC;
+ const uint16_t* B1 = B0 + a.bK * F;
+ const uint16_t* A1 = A0 + 1 * dA;
+ const uint16_t* A2 = A0 + 2 * dA;
+ const uint16_t* A3 = A0 + 3 * dA;
+ const uint16_t* A4 = A0 + 4 * dA;
+ if (N > F)
+ {
+ if (update)
+ {
+ if (M > 0) c00 = svld1_f32(body, C + 0 * dC + 0), c01 = svld1_f32(body, C + 0 * dC + F);
+ if (M > 1) c10 = svld1_f32(body, C + 1 * dC + 0), c11 = svld1_f32(body, C + 1 * dC + F);
+ if (M > 2) c20 = svld1_f32(body, C + 2 * dC + 0), c21 = svld1_f32(body, C + 2 * dC + F);
+ if (M > 3) c30 = svld1_f32(body, C + 3 * dC + 0), c31 = svld1_f32(body, C + 3 * dC + F);
+ if (M > 4) c40 = svld1_f32(body, C + 4 * dC + 0), c41 = svld1_f32(body, C + 4 * dC + F);
+ }
+ else
+ {
+ if (M > 0) c00 = svdup_n_f32(0.0f), c01 = svdup_n_f32(0.0f);
+ if (M > 1) c10 = svdup_n_f32(0.0f), c11 = svdup_n_f32(0.0f);
+ if (M > 2) c20 = svdup_n_f32(0.0f), c21 = svdup_n_f32(0.0f);
+ if (M > 3) c30 = svdup_n_f32(0.0f), c31 = svdup_n_f32(0.0f);
+ if (M > 4) c40 = svdup_n_f32(0.0f), c41 = svdup_n_f32(0.0f);
+ }
+ for (size_t k = 0; k < K; k += 2)
+ {
+ svuint32_t b0u = svld1_u32(body, (const uint32_t*)B0);
+ b00 = svreinterpret_f32_u32(svlsl_n_u32_x(body, b0u, Base::Bf16::SHIFT));
+ b01 = svreinterpret_f32_u32(svand_u32_x(body, b0u, maskBits));
+ svuint32_t b1u = svld1_u32(body, (const uint32_t*)B1);
+ b10 = svreinterpret_f32_u32(svlsl_n_u32_x(body, b1u, Base::Bf16::SHIFT));
+ b11 = svreinterpret_f32_u32(svand_u32_x(body, b1u, maskBits));
+ if (M > 0)
+ {
+ a0 = BroadcastBf16(A0[k + 0]);
+ c00 = svmla_f32_x(body, c00, a0, b00);
+ c01 = svmla_f32_x(body, c01, a0, b10);
+ a0 = BroadcastBf16(A0[k + 1]);
+ c00 = svmla_f32_x(body, c00, a0, b01);
+ c01 = svmla_f32_x(body, c01, a0, b11);
+ }
+ if (M > 1)
+ {
+ a0 = BroadcastBf16(A1[k + 0]);
+ c10 = svmla_f32_x(body, c10, a0, b00);
+ c11 = svmla_f32_x(body, c11, a0, b10);
+ a0 = BroadcastBf16(A1[k + 1]);
+ c10 = svmla_f32_x(body, c10, a0, b01);
+ c11 = svmla_f32_x(body, c11, a0, b11);
+ }
+ if (M > 2)
+ {
+ a0 = BroadcastBf16(A2[k + 0]);
+ c20 = svmla_f32_x(body, c20, a0, b00);
+ c21 = svmla_f32_x(body, c21, a0, b10);
+ a0 = BroadcastBf16(A2[k + 1]);
+ c20 = svmla_f32_x(body, c20, a0, b01);
+ c21 = svmla_f32_x(body, c21, a0, b11);
+ }
+ if (M > 3)
+ {
+ a0 = BroadcastBf16(A3[k + 0]);
+ c30 = svmla_f32_x(body, c30, a0, b00);
+ c31 = svmla_f32_x(body, c31, a0, b10);
+ a0 = BroadcastBf16(A3[k + 1]);
+ c30 = svmla_f32_x(body, c30, a0, b01);
+ c31 = svmla_f32_x(body, c31, a0, b11);
+ }
+ if (M > 4)
+ {
+ a0 = BroadcastBf16(A4[k + 0]);
+ c40 = svmla_f32_x(body, c40, a0, b00);
+ c41 = svmla_f32_x(body, c41, a0, b10);
+ a0 = BroadcastBf16(A4[k + 1]);
+ c40 = svmla_f32_x(body, c40, a0, b01);
+ c41 = svmla_f32_x(body, c41, a0, b11);
+ }
+ B0 += DF;
+ B1 += DF;
+ }
+ svbool_t mask1 = (N == DF) ? body : svwhilelt_b32((size_t)0, N - F);
+ if (M > 0) Save2(dst, C + 0 * dC, c00, c01, bias0, bias1, param0, param1, body, mask1, F), C += dC, dst += dD;
+ if (M > 1) Save2(dst, C + 0 * dC, c10, c11, bias0, bias1, param0, param1, body, mask1, F), C += dC, dst += dD;
+ if (M > 2) Save2(dst, C + 0 * dC, c20, c21, bias0, bias1, param0, param1, body, mask1, F), C += dC, dst += dD;
+ if (M > 3) Save2(dst, C + 0 * dC, c30, c31, bias0, bias1, param0, param1, body, mask1, F), C += dC, dst += dD;
+ if (M > 4) Save2(dst, C + 0 * dC, c40, c41, bias0, bias1, param0, param1, body, mask1, F), C += dC, dst += dD;
+ }
+ else
+ {
+ if (update)
+ {
+ if (M > 0) c00 = svld1_f32(body, C + 0 * dC + 0);
+ if (M > 1) c10 = svld1_f32(body, C + 1 * dC + 0);
+ if (M > 2) c20 = svld1_f32(body, C + 2 * dC + 0);
+ if (M > 3) c30 = svld1_f32(body, C + 3 * dC + 0);
+ if (M > 4) c40 = svld1_f32(body, C + 4 * dC + 0);
+ }
+ else
+ {
+ if (M > 0) c00 = svdup_n_f32(0.0f);
+ if (M > 1) c10 = svdup_n_f32(0.0f);
+ if (M > 2) c20 = svdup_n_f32(0.0f);
+ if (M > 3) c30 = svdup_n_f32(0.0f);
+ if (M > 4) c40 = svdup_n_f32(0.0f);
+ }
+ for (size_t k = 0; k < K; k += 2)
+ {
+ svuint32_t b0u = svld1_u32(body, (const uint32_t*)B0);
+ b00 = svreinterpret_f32_u32(svlsl_n_u32_x(body, b0u, Base::Bf16::SHIFT));
+ b01 = svreinterpret_f32_u32(svand_u32_x(body, b0u, maskBits));
+ if (M > 0)
+ {
+ a0 = BroadcastBf16(A0[k + 0]);
+ c00 = svmla_f32_x(body, c00, a0, b00);
+ a0 = BroadcastBf16(A0[k + 1]);
+ c00 = svmla_f32_x(body, c00, a0, b01);
+ }
+ if (M > 1)
+ {
+ a0 = BroadcastBf16(A1[k + 0]);
+ c10 = svmla_f32_x(body, c10, a0, b00);
+ a0 = BroadcastBf16(A1[k + 1]);
+ c10 = svmla_f32_x(body, c10, a0, b01);
+ }
+ if (M > 2)
+ {
+ a0 = BroadcastBf16(A2[k + 0]);
+ c20 = svmla_f32_x(body, c20, a0, b00);
+ a0 = BroadcastBf16(A2[k + 1]);
+ c20 = svmla_f32_x(body, c20, a0, b01);
+ }
+ if (M > 3)
+ {
+ a0 = BroadcastBf16(A3[k + 0]);
+ c30 = svmla_f32_x(body, c30, a0, b00);
+ a0 = BroadcastBf16(A3[k + 1]);
+ c30 = svmla_f32_x(body, c30, a0, b01);
+ }
+ if (M > 4)
+ {
+ a0 = BroadcastBf16(A4[k + 0]);
+ c40 = svmla_f32_x(body, c40, a0, b00);
+ a0 = BroadcastBf16(A4[k + 1]);
+ c40 = svmla_f32_x(body, c40, a0, b01);
+ }
+ B0 += DF;
+ }
+ svbool_t mask0 = (N == F) ? body : svwhilelt_b32((size_t)0, N);
+ if (M > 0) Save1(dst, C + 0 * dC, c00, bias0, param0, param1, 0, mask0), C += dC, dst += dD;
+ if (M > 1) Save1(dst, C + 0 * dC, c10, bias0, param0, param1, 0, mask0), C += dC, dst += dD;
+ if (M > 2) Save1(dst, C + 0 * dC, c20, bias0, param0, param1, 0, mask0), C += dC, dst += dD;
+ if (M > 3) Save1(dst, C + 0 * dC, c30, bias0, param0, param1, 0, mask0), C += dC, dst += dD;
+ if (M > 4) Save1(dst, C + 0 * dC, c40, bias0, param0, param1, 0, mask0), C += dC, dst += dD;
+ }
+ }
+
+ typedef void(*GemmNN_2xM_Ptr)(const uint16_t* A0, const InnerProductParam16b& p, const AlgParam& a,
+ size_t N, size_t K, int update, const uint16_t* B0, float* C,
+ svfloat32_t bias0, svfloat32_t bias1, svfloat32_t param0, svfloat32_t param1, uint8_t* dst);
+
+ template GemmNN_2xM_Ptr GetGemmNN_2xM(size_t M)
+ {
+ switch (M)
+ {
+ case 0: return NULL;
+ case 1: return InnerProduct16bGemmNN_2xM;
+ case 2: return InnerProduct16bGemmNN_2xM;
+ case 3: return InnerProduct16bGemmNN_2xM;
+ case 4: return InnerProduct16bGemmNN_2xM;
+ case 5: return InnerProduct16bGemmNN_2xM;
+ }
+ assert(0);
+ return NULL;
+ }
+
+ template void InnerProduct16bGemmNN_Gemm2(
+ const uint16_t* A, const InnerProductParam16b& p, const AlgParam& a,
+ size_t M, size_t N, size_t K, int update, const uint16_t* B, float* C, int post,
+ const float* bias, const float* params, float* sum, uint8_t* dst)
+ {
+ const size_t F = a.F, DF = F * 2;
+ const svbool_t body = svptrue_b32();
+ size_t m1 = M, m = 5;
+ size_t mm = AlignLoAny(m1, m), t = m1 - mm;
+ size_t dA = a.aK, dB = a.bK * DF, dC = a.cN, dD = p.N * a.eC;
+ GemmNN_2xM_Ptr gemm_2xM = post ? GetGemmNN_2xM(m) : GetGemmNN_2xM(m);
+ GemmNN_2xM_Ptr gemm_2xT = post ? GetGemmNN_2xM(t) : GetGemmNN_2xM(t);
+
+ svfloat32_t _param0 = svdup_n_f32(params[0]);
+ svfloat32_t _param1 = svdup_n_f32(params[1]);
+ for (size_t j = 0; j < N; j += DF)
+ {
+ size_t dN = Simd::Min(DF, N - j);
+ svfloat32_t _bias0 = svld1_f32(body, bias + j + 0);
+ svfloat32_t _bias1 = svld1_f32(body, bias + j + F);
+ if (type == ::SimdConvolutionActivationPrelu)
+ {
+ _param0 = svld1_f32(body, params + j + 0);
+ _param1 = svld1_f32(body, params + j + F);
+ }
+
+ size_t i = 0;
+ for (; i < mm; i += m)
+ gemm_2xM(A + i * dA, p, a, dN, K, update, B, C + i * dC, _bias0, _bias1, _param0, _param1, dst + i * dD);
+ for (; i < m1; i += t)
+ gemm_2xT(A + i * dA, p, a, dN, K, update, B, C + i * dC, _bias0, _bias1, _param0, _param1, dst + i * dD);
+ B += dB;
+ C += dN;
+ dst += DF * a.eC;
+ }
+ }
+
+ //-------------------------------------------------------------------------------------------------
+
+ template SIMD_INLINE void SetGemm(const InnerProductParam16b& p, GemmPtr& gemm)
+ {
+ if (p.typeC == SimdTensorData16b)
+ gemm = InnerProduct16bGemmNN_Gemm2;
+ else
+ gemm = InnerProduct16bGemmNN_Gemm2;
+ }
+
+ SynetInnerProduct16bGemmNN::SynetInnerProduct16bGemmNN(const InnerProductParam16b& p)
+ : Base::SynetInnerProduct16bGemmNN(p)
+ {
+ const size_t F = svcntw();
+ SetAlgParam(F, 5, F * 2, 2, Base::AlgCacheL1(), Base::AlgCacheL2(), Base::AlgCacheL3());
+ if (_sizeA)
+ {
+ if (p.typeA == SimdTensorData16b)
+ _prepA = InnerProduct16bGemmNN_ReorderA;
+ else
+ _prepA = InnerProduct16bGemmNN_ConvertA;
+ }
+ if (p.typeB == SimdTensorData32f || p.constB)
+ {
+ if (p.transB)
+ _prepB = InnerProduct16bGemmNN_ConvertBt;
+ else
+ _prepB = InnerProduct16bGemmNN_ConvertBn;
+ }
+ else
+ {
+ if (p.transB)
+ _prepB = InnerProduct16bGemmNN_ReorderBt;
+ else
+ _prepB = InnerProduct16bGemmNN_ReorderBn;
+ }
+ switch (p.activation)
+ {
+ case SimdConvolutionActivationIdentity: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationRelu: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationLeakyRelu: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationRestrictRange: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationPrelu: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationElu: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationHswish: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationMish: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationHardSigmoid: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationSwish: SetGemm(p, _gemm); break;
+ case SimdConvolutionActivationGelu: SetGemm(p, _gemm); break;
+ default: assert(0);
+ }
+ }
+ }
+#endif
+}
diff --git a/src/Simd/SimdSynetInnerProduct16b.h b/src/Simd/SimdSynetInnerProduct16b.h
index e88a1e0d9d..c860ef39fe 100644
--- a/src/Simd/SimdSynetInnerProduct16b.h
+++ b/src/Simd/SimdSynetInnerProduct16b.h
@@ -243,6 +243,23 @@ namespace Simd
void* SynetInnerProduct16bInit(size_t M, size_t N, size_t K, SimdTensorDataType typeA, SimdTensorDataType typeB, SimdTensorDataType typeC, SimdBool transB, SimdBool constB, SimdBool bias, SimdConvolutionActivationType activation);
}
#endif
+
+#ifdef SIMD_SVE2_ENABLE
+ namespace Sve2
+ {
+ class SynetInnerProduct16bGemmNN : public Base::SynetInnerProduct16bGemmNN
+ {
+ public:
+ SynetInnerProduct16bGemmNN(const InnerProductParam16b& p);
+
+ virtual String Ext() const { return "Sve2"; }
+ };
+
+ //-------------------------------------------------------------------------------------------------
+
+ void* SynetInnerProduct16bInit(size_t M, size_t N, size_t K, SimdTensorDataType typeA, SimdTensorDataType typeB, SimdTensorDataType typeC, SimdBool transB, SimdBool constB, SimdBool bias, SimdConvolutionActivationType activation);
+ }
+#endif
}
#endif
diff --git a/src/Test/TestSynetInnerProduct16b.cpp b/src/Test/TestSynetInnerProduct16b.cpp
index 6ac6582d2e..dbaaba2d5f 100644
--- a/src/Test/TestSynetInnerProduct16b.cpp
+++ b/src/Test/TestSynetInnerProduct16b.cpp
@@ -280,6 +280,11 @@ namespace Test
result = result && SynetInnerProduct16bForwardAutoTest(EPS, FUNC_IP16B(Simd::AmxBf16::SynetInnerProduct16bInit), FUNC_IP16B(SimdSynetInnerProduct16bInit));
#endif
+#ifdef SIMD_SVE2_ENABLE
+ if (Simd::Sve2::Enable && TestSve2(options))
+ result = result && SynetInnerProduct16bForwardAutoTest(EPS, FUNC_IP16B(Simd::Sve2::SynetInnerProduct16bInit), FUNC_IP16B(SimdSynetInnerProduct16bInit));
+#endif
+
return result;
}
#endif