Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/2026.html
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ <h5>New features</h5>
<li>SVE2 optimizations of class SynetConvolution32fDepthwiseDotProduct.</li>
<li>SVE2 optimizations of class SynetConvolution32fNhwcDirect.</li>
<li>SVE2 optimizations of class SynetInnerProduct32fGemm.</li>
<li>SVE2 optimizations of class SynetInnerProduct32fProd.</li>
</ul>
<h5>Renaming</h5>
<ul>
Expand Down
333 changes: 332 additions & 1 deletion src/Simd/SimdSve2SynetInnerProduct32f.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -61,12 +61,343 @@ namespace Simd

//-----------------------------------------------------------------------------------------

void InnerProductKxKNr1x1(size_t K, const float* src, const float* weight0, const float* bias, float* dst, const svbool_t& tail)
{
const size_t F = svcntw();
const svbool_t body = svptrue_b32();
svfloat32_t d00 = svld1_f32(body, bias + 0 * F);
svfloat32_t s0, s1, s2, s3, w0, w1, w2, w3;
size_t K2 = AlignLo(K, 2);
size_t K4 = AlignLo(K, 4);
size_t k = 0, off = 0;
for (; k < K4; k += 4, off += F * 4)
{
s0 = svdup_n_f32(src[k + 0]);
s1 = svdup_n_f32(src[k + 1]);
s2 = svdup_n_f32(src[k + 2]);
s3 = svdup_n_f32(src[k + 3]);
w0 = svld1_f32(body, weight0 + off + 0 * F);
w1 = svld1_f32(body, weight0 + off + 1 * F);
w2 = svld1_f32(body, weight0 + off + 2 * F);
w3 = svld1_f32(body, weight0 + off + 3 * F);
d00 = svmla_f32_x(body, d00, w0, s0);
d00 = svmla_f32_x(body, d00, w1, s1);
d00 = svmla_f32_x(body, d00, w2, s2);
d00 = svmla_f32_x(body, d00, w3, s3);
}
for (; k < K2; k += 2, off += F * 2)
{
s0 = svdup_n_f32(src[k + 0]);
s1 = svdup_n_f32(src[k + 1]);
w0 = svld1_f32(body, weight0 + off + 0 * F);
w1 = svld1_f32(body, weight0 + off + 1 * F);
d00 = svmla_f32_x(body, d00, w0, s0);
d00 = svmla_f32_x(body, d00, w1, s1);
}
for (; k < K; k++, off += F)
{
s0 = svdup_n_f32(src[k]);
w0 = svld1_f32(body, weight0 + off);
d00 = svmla_f32_x(body, d00, w0, s0);
}
svst1_f32(tail, dst + 0 * F, d00);
}

void InnerProductKxKNr1x4(size_t K, const float* src, const float* weight0, const float* bias, float* dst)
{
const size_t F = svcntw();
const svbool_t body = svptrue_b32();
svfloat32_t d00 = svld1_f32(body, bias + 0 * F);
svfloat32_t d01 = svld1_f32(body, bias + 1 * F);
svfloat32_t d02 = svld1_f32(body, bias + 2 * F);
svfloat32_t d03 = svld1_f32(body, bias + 3 * F);
svfloat32_t s0, s1, s2, s3, w00, w01, w02, w03, w10, w11, w12, w13;
const float* weight1 = weight0 + 1 * K * F;
const float* weight2 = weight0 + 2 * K * F;
const float* weight3 = weight0 + 3 * K * F;
size_t K2 = AlignLo(K, 2);
size_t K4 = AlignLo(K, 4);
size_t k = 0, off = 0;
for (; k < K4; k += 4, off += F * 4)
{
s0 = svdup_n_f32(src[k + 0]);
s1 = svdup_n_f32(src[k + 1]);
s2 = svdup_n_f32(src[k + 2]);
s3 = svdup_n_f32(src[k + 3]);
w00 = svld1_f32(body, weight0 + off + 0 * F);
w01 = svld1_f32(body, weight0 + off + 1 * F);
w02 = svld1_f32(body, weight0 + off + 2 * F);
w03 = svld1_f32(body, weight0 + off + 3 * F);
w10 = svld1_f32(body, weight1 + off + 0 * F);
w11 = svld1_f32(body, weight1 + off + 1 * F);
w12 = svld1_f32(body, weight1 + off + 2 * F);
w13 = svld1_f32(body, weight1 + off + 3 * F);
d00 = svmla_f32_x(body, d00, w00, s0);
d01 = svmla_f32_x(body, d01, w10, s0);
d00 = svmla_f32_x(body, d00, w01, s1);
d01 = svmla_f32_x(body, d01, w11, s1);
d00 = svmla_f32_x(body, d00, w02, s2);
d01 = svmla_f32_x(body, d01, w12, s2);
d00 = svmla_f32_x(body, d00, w03, s3);
d01 = svmla_f32_x(body, d01, w13, s3);
w00 = svld1_f32(body, weight2 + off + 0 * F);
w01 = svld1_f32(body, weight2 + off + 1 * F);
w02 = svld1_f32(body, weight2 + off + 2 * F);
w03 = svld1_f32(body, weight2 + off + 3 * F);
w10 = svld1_f32(body, weight3 + off + 0 * F);
w11 = svld1_f32(body, weight3 + off + 1 * F);
w12 = svld1_f32(body, weight3 + off + 2 * F);
w13 = svld1_f32(body, weight3 + off + 3 * F);
d02 = svmla_f32_x(body, d02, w00, s0);
d03 = svmla_f32_x(body, d03, w10, s0);
d02 = svmla_f32_x(body, d02, w01, s1);
d03 = svmla_f32_x(body, d03, w11, s1);
d02 = svmla_f32_x(body, d02, w02, s2);
d03 = svmla_f32_x(body, d03, w12, s2);
d02 = svmla_f32_x(body, d02, w03, s3);
d03 = svmla_f32_x(body, d03, w13, s3);
}
for (; k < K2; k += 2, off += F * 2)
{
s0 = svdup_n_f32(src[k + 0]);
s1 = svdup_n_f32(src[k + 1]);
w00 = svld1_f32(body, weight0 + off + 0 * F);
w01 = svld1_f32(body, weight0 + off + 1 * F);
w10 = svld1_f32(body, weight1 + off + 0 * F);
w11 = svld1_f32(body, weight1 + off + 1 * F);
d00 = svmla_f32_x(body, d00, w00, s0);
d01 = svmla_f32_x(body, d01, w10, s0);
d00 = svmla_f32_x(body, d00, w01, s1);
d01 = svmla_f32_x(body, d01, w11, s1);
w00 = svld1_f32(body, weight2 + off + 0 * F);
w01 = svld1_f32(body, weight2 + off + 1 * F);
w10 = svld1_f32(body, weight3 + off + 0 * F);
w11 = svld1_f32(body, weight3 + off + 1 * F);
d02 = svmla_f32_x(body, d02, w00, s0);
d03 = svmla_f32_x(body, d03, w10, s0);
d02 = svmla_f32_x(body, d02, w01, s1);
d03 = svmla_f32_x(body, d03, w11, s1);
}
for (; k < K; k++, off += F)
{
s0 = svdup_n_f32(src[k + 0]);
w00 = svld1_f32(body, weight0 + off + 0 * F);
w10 = svld1_f32(body, weight1 + off + 0 * F);
d00 = svmla_f32_x(body, d00, w00, s0);
d01 = svmla_f32_x(body, d01, w10, s0);
w00 = svld1_f32(body, weight2 + off + 0 * F);
w10 = svld1_f32(body, weight3 + off + 0 * F);
d02 = svmla_f32_x(body, d02, w00, s0);
d03 = svmla_f32_x(body, d03, w10, s0);
}
svst1_f32(body, dst + 0 * F, d00);
svst1_f32(body, dst + 1 * F, d01);
svst1_f32(body, dst + 2 * F, d02);
svst1_f32(body, dst + 3 * F, d03);
}

void InnerProductKxKNr1x8(size_t K, const float* src, const float* weight0, const float* bias, float* dst)
{
const size_t F = svcntw();
const svbool_t body = svptrue_b32();
svfloat32_t d00 = svld1_f32(body, bias + 0 * F);
svfloat32_t d01 = svld1_f32(body, bias + 1 * F);
svfloat32_t d02 = svld1_f32(body, bias + 2 * F);
svfloat32_t d03 = svld1_f32(body, bias + 3 * F);
svfloat32_t d04 = svld1_f32(body, bias + 4 * F);
svfloat32_t d05 = svld1_f32(body, bias + 5 * F);
svfloat32_t d06 = svld1_f32(body, bias + 6 * F);
svfloat32_t d07 = svld1_f32(body, bias + 7 * F);
svfloat32_t s0, s1, s2, s3, w00, w01, w10, w11;
const float* weight1 = weight0 + 1 * K * F;
const float* weight2 = weight0 + 2 * K * F;
const float* weight3 = weight0 + 3 * K * F;
size_t K2 = AlignLo(K, 2);
size_t K4 = AlignLo(K, 4);
size_t k = 0, off0 = 0, off4 = 4 * K * F;
for (; k < K4; k += 4, off0 += F * 4, off4 += F * 4)
{
s0 = svdup_n_f32(src[k + 0]);
s1 = svdup_n_f32(src[k + 1]);
s2 = svdup_n_f32(src[k + 2]);
s3 = svdup_n_f32(src[k + 3]);

w00 = svld1_f32(body, weight0 + off0 + 0 * F);
w01 = svld1_f32(body, weight0 + off0 + 1 * F);
w10 = svld1_f32(body, weight1 + off0 + 0 * F);
w11 = svld1_f32(body, weight1 + off0 + 1 * F);
d00 = svmla_f32_x(body, d00, w00, s0);
d01 = svmla_f32_x(body, d01, w10, s0);
d00 = svmla_f32_x(body, d00, w01, s1);
d01 = svmla_f32_x(body, d01, w11, s1);
w00 = svld1_f32(body, weight0 + off0 + 2 * F);
w01 = svld1_f32(body, weight0 + off0 + 3 * F);
w10 = svld1_f32(body, weight1 + off0 + 2 * F);
w11 = svld1_f32(body, weight1 + off0 + 3 * F);
d00 = svmla_f32_x(body, d00, w00, s2);
d01 = svmla_f32_x(body, d01, w10, s2);
d00 = svmla_f32_x(body, d00, w01, s3);
d01 = svmla_f32_x(body, d01, w11, s3);
w00 = svld1_f32(body, weight2 + off0 + 0 * F);
w01 = svld1_f32(body, weight2 + off0 + 1 * F);
w10 = svld1_f32(body, weight3 + off0 + 0 * F);
w11 = svld1_f32(body, weight3 + off0 + 1 * F);
d02 = svmla_f32_x(body, d02, w00, s0);
d03 = svmla_f32_x(body, d03, w10, s0);
d02 = svmla_f32_x(body, d02, w01, s1);
d03 = svmla_f32_x(body, d03, w11, s1);
w00 = svld1_f32(body, weight2 + off0 + 2 * F);
w01 = svld1_f32(body, weight2 + off0 + 3 * F);
w10 = svld1_f32(body, weight3 + off0 + 2 * F);
w11 = svld1_f32(body, weight3 + off0 + 3 * F);
d02 = svmla_f32_x(body, d02, w00, s2);
d03 = svmla_f32_x(body, d03, w10, s2);
d02 = svmla_f32_x(body, d02, w01, s3);
d03 = svmla_f32_x(body, d03, w11, s3);

w00 = svld1_f32(body, weight0 + off4 + 0 * F);
w01 = svld1_f32(body, weight0 + off4 + 1 * F);
w10 = svld1_f32(body, weight1 + off4 + 0 * F);
w11 = svld1_f32(body, weight1 + off4 + 1 * F);
d04 = svmla_f32_x(body, d04, w00, s0);
d05 = svmla_f32_x(body, d05, w10, s0);
d04 = svmla_f32_x(body, d04, w01, s1);
d05 = svmla_f32_x(body, d05, w11, s1);
w00 = svld1_f32(body, weight0 + off4 + 2 * F);
w01 = svld1_f32(body, weight0 + off4 + 3 * F);
w10 = svld1_f32(body, weight1 + off4 + 2 * F);
w11 = svld1_f32(body, weight1 + off4 + 3 * F);
d04 = svmla_f32_x(body, d04, w00, s2);
d05 = svmla_f32_x(body, d05, w10, s2);
d04 = svmla_f32_x(body, d04, w01, s3);
d05 = svmla_f32_x(body, d05, w11, s3);
w00 = svld1_f32(body, weight2 + off4 + 0 * F);
w01 = svld1_f32(body, weight2 + off4 + 1 * F);
w10 = svld1_f32(body, weight3 + off4 + 0 * F);
w11 = svld1_f32(body, weight3 + off4 + 1 * F);
d06 = svmla_f32_x(body, d06, w00, s0);
d07 = svmla_f32_x(body, d07, w10, s0);
d06 = svmla_f32_x(body, d06, w01, s1);
d07 = svmla_f32_x(body, d07, w11, s1);
w00 = svld1_f32(body, weight2 + off4 + 2 * F);
w01 = svld1_f32(body, weight2 + off4 + 3 * F);
w10 = svld1_f32(body, weight3 + off4 + 2 * F);
w11 = svld1_f32(body, weight3 + off4 + 3 * F);
d06 = svmla_f32_x(body, d06, w00, s2);
d07 = svmla_f32_x(body, d07, w10, s2);
d06 = svmla_f32_x(body, d06, w01, s3);
d07 = svmla_f32_x(body, d07, w11, s3);
}
for (; k < K2; k += 2, off0 += F * 2, off4 += F * 2)
{
s0 = svdup_n_f32(src[k + 0]);
s1 = svdup_n_f32(src[k + 1]);

w00 = svld1_f32(body, weight0 + off0 + 0 * F);
w01 = svld1_f32(body, weight0 + off0 + 1 * F);
w10 = svld1_f32(body, weight1 + off0 + 0 * F);
w11 = svld1_f32(body, weight1 + off0 + 1 * F);
d00 = svmla_f32_x(body, d00, w00, s0);
d01 = svmla_f32_x(body, d01, w10, s0);
d00 = svmla_f32_x(body, d00, w01, s1);
d01 = svmla_f32_x(body, d01, w11, s1);
w00 = svld1_f32(body, weight2 + off0 + 0 * F);
w01 = svld1_f32(body, weight2 + off0 + 1 * F);
w10 = svld1_f32(body, weight3 + off0 + 0 * F);
w11 = svld1_f32(body, weight3 + off0 + 1 * F);
d02 = svmla_f32_x(body, d02, w00, s0);
d03 = svmla_f32_x(body, d03, w10, s0);
d02 = svmla_f32_x(body, d02, w01, s1);
d03 = svmla_f32_x(body, d03, w11, s1);

w00 = svld1_f32(body, weight0 + off4 + 0 * F);
w01 = svld1_f32(body, weight0 + off4 + 1 * F);
w10 = svld1_f32(body, weight1 + off4 + 0 * F);
w11 = svld1_f32(body, weight1 + off4 + 1 * F);
d04 = svmla_f32_x(body, d04, w00, s0);
d05 = svmla_f32_x(body, d05, w10, s0);
d04 = svmla_f32_x(body, d04, w01, s1);
d05 = svmla_f32_x(body, d05, w11, s1);
w00 = svld1_f32(body, weight2 + off4 + 0 * F);
w01 = svld1_f32(body, weight2 + off4 + 1 * F);
w10 = svld1_f32(body, weight3 + off4 + 0 * F);
w11 = svld1_f32(body, weight3 + off4 + 1 * F);
d06 = svmla_f32_x(body, d06, w00, s0);
d07 = svmla_f32_x(body, d07, w10, s0);
d06 = svmla_f32_x(body, d06, w01, s1);
d07 = svmla_f32_x(body, d07, w11, s1);
}
for (; k < K; k++, off0 += F, off4 += F)
{
s0 = svdup_n_f32(src[k + 0]);

w00 = svld1_f32(body, weight0 + off0 + 0 * F);
w10 = svld1_f32(body, weight1 + off0 + 0 * F);
d00 = svmla_f32_x(body, d00, w00, s0);
d01 = svmla_f32_x(body, d01, w10, s0);
w00 = svld1_f32(body, weight2 + off0 + 0 * F);
w10 = svld1_f32(body, weight3 + off0 + 0 * F);
d02 = svmla_f32_x(body, d02, w00, s0);
d03 = svmla_f32_x(body, d03, w10, s0);

w00 = svld1_f32(body, weight0 + off4 + 0 * F);
w10 = svld1_f32(body, weight1 + off4 + 0 * F);
d04 = svmla_f32_x(body, d04, w00, s0);
d05 = svmla_f32_x(body, d05, w10, s0);
w00 = svld1_f32(body, weight2 + off4 + 0 * F);
w10 = svld1_f32(body, weight3 + off4 + 0 * F);
d06 = svmla_f32_x(body, d06, w00, s0);
d07 = svmla_f32_x(body, d07, w10, s0);
}
svst1_f32(body, dst + 0 * F, d00);
svst1_f32(body, dst + 1 * F, d01);
svst1_f32(body, dst + 2 * F, d02);
svst1_f32(body, dst + 3 * F, d03);
svst1_f32(body, dst + 4 * F, d04);
svst1_f32(body, dst + 5 * F, d05);
svst1_f32(body, dst + 6 * F, d06);
svst1_f32(body, dst + 7 * F, d07);
}

void InnerProductKxKNr(const float* src, const float* weight, const float* bias, size_t input, size_t output, float* dst)
{
const size_t F = svcntw();
const svbool_t body = svptrue_b32();
size_t outputF1 = AlignLo(output, F * 1);
size_t outputF4 = AlignLo(output, F * 4);
size_t outputF8 = AlignLo(output, F * 8);
size_t o = 0;
for (; o < outputF8; o += F * 8)
InnerProductKxKNr1x8(input, src, weight + o * input, bias + o, dst + o);
for (; o < outputF4; o += F * 4)
InnerProductKxKNr1x4(input, src, weight + o * input, bias + o, dst + o);
for (; o < outputF1; o += F * 1)
InnerProductKxKNr1x1(input, src, weight + o * input, bias + o, dst + o, body);
if (o < output)
InnerProductKxKNr1x1(input, src, weight + o * input, bias + o, dst + o, svwhilelt_b32(o, output));
}

SynetInnerProduct32fProd::SynetInnerProduct32fProd(const InnerProductParam32f& p)
: Base::SynetInnerProduct32fProd(p)
{
if (_param.N > 1)
{
SetSize(svcntw());
_prod = InnerProductKxKNr;
}
}

//-----------------------------------------------------------------------------------------

void* SynetInnerProduct32fInit(size_t M, size_t N, size_t K, SimdBool transB, SimdBool constB, SimdBool bias, SimdConvolutionActivationType activation)
{
InnerProductParam32f param(M, N, K, transB, constB, bias, activation);
if (!param.Valid())
return NULL;
return new SynetInnerProduct32fGemm(param);
if (SynetInnerProduct32fProd::Preferable(param))
return new SynetInnerProduct32fProd(param);
else
return new SynetInnerProduct32fGemm(param);
}
}
#endif
Expand Down
8 changes: 8 additions & 0 deletions src/Simd/SimdSynetInnerProduct32f.h
Original file line number Diff line number Diff line change
Expand Up @@ -278,6 +278,14 @@ namespace Simd
virtual String Ext() const { return "Sve2"; }
};

class SynetInnerProduct32fProd : public Base::SynetInnerProduct32fProd
{
public:
SynetInnerProduct32fProd(const InnerProductParam32f& p);

virtual String Ext() const { return "Sve2"; }
};

void* SynetInnerProduct32fInit(size_t M, size_t N, size_t K, SimdBool transB, SimdBool constB, SimdBool bias, SimdConvolutionActivationType activation);
}
#endif
Expand Down