1
0
mirror of https://github.com/opencv/opencv.git synced 2026-07-21 19:33:03 +04:00

Merge pull request #29242 from amd:fast_gemm_simd

Optimized gemm implementation with Universal SIMD #29242

- vectorized GEMMSingleMul and GEMMBlockMul

### Pull Request Readiness Checklist

See details at https://github.com/opencv/opencv/wiki/How_to_contribute#making-a-good-pull-request

- [x] I agree to contribute to the project under Apache 2 License.
- [x] To the best of my knowledge, the proposed patch is not based on a code under GPL or another license that is incompatible with OpenCV
- [x] The PR is proposed to the proper branch
- [ ] There is a reference to the original bug report and related work
- [x] There is accuracy test, performance test and test data in opencv_extra repository, if applicable
      Patch to opencv_extra has the same branch name.
- [x] The feature is well documented and sample code can be built with the project CMake
This commit is contained in:
Madan mohan Manokar
2026-07-01 16:17:43 +05:30
committed by GitHub
parent 5b83b0c440
commit 28a1d0dddb
+574 -45
View File
@@ -13,6 +13,7 @@
// Copyright (C) 2000-2008, Intel Corporation, all rights reserved.
// Copyright (C) 2009-2011, Willow Garage Inc., all rights reserved.
// Copyright (C) 2014-2015, Itseez Inc., all rights reserved.
// Copyright (C) 2026, Advanced Micro Devices, Inc., all rights reserved.
// Third party copyrights are property of their respective owners.
//
// Redistribution and use in source and binary forms, with or without modification,
@@ -175,6 +176,337 @@ GEMM_TransposeBlock( const uchar* src, size_t src_step,
}
}
#if CV_SIMD
template<typename ST> struct VecTraits;
template<> struct VecTraits<float> {
typedef v_float32 VT;
static inline VT zero() { return vx_setzero_f32(); }
static inline VT load(const float* p) { return vx_load(p); }
static inline VT setall(float v) { return vx_setall_f32(v); }
};
#endif
#if CV_SIMD_64F
template<> struct VecTraits<double> {
typedef v_float64 VT;
static inline VT zero() { return vx_setzero_f64(); }
static inline VT load(const double* p) { return vx_load(p); }
static inline VT setall(double v) { return vx_setall_f64(v); }
};
#endif
#if CV_SIMD
template<typename ST>
static inline ST simdDotProduct(const ST* a, const ST* b, int n, int& k)
{
typedef typename VecTraits<ST>::VT VecT;
const int vlanes = VTraits<VecT>::vlanes();
const int stride4 = vlanes * 4;
VecT vs0 = VecTraits<ST>::zero();
VecT vs1 = VecTraits<ST>::zero();
VecT vs2 = VecTraits<ST>::zero();
VecT vs3 = VecTraits<ST>::zero();
for( ; k <= n - stride4; k += stride4 )
{
vs0 = v_muladd(VecTraits<ST>::load(a + k), VecTraits<ST>::load(b + k), vs0);
vs1 = v_muladd(VecTraits<ST>::load(a + k + vlanes), VecTraits<ST>::load(b + k + vlanes), vs1);
vs2 = v_muladd(VecTraits<ST>::load(a + k + vlanes * 2), VecTraits<ST>::load(b + k + vlanes * 2), vs2);
vs3 = v_muladd(VecTraits<ST>::load(a + k + vlanes * 3), VecTraits<ST>::load(b + k + vlanes * 3), vs3);
}
vs0 = v_add(v_add(vs0, vs1), v_add(vs2, vs3));
for( ; k <= n - vlanes; k += vlanes )
vs0 = v_muladd(VecTraits<ST>::load(a + k), VecTraits<ST>::load(b + k), vs0);
return v_reduce_sum(vs0);
}
#endif
#if CV_SIMD_64F
// Float input, double accumulation dot product (for 32FC1 where WT=double)
static inline double simdDotProduct_f32f64(const float* a, const float* b, int n, int& k)
{
const int vlanes_f32 = VTraits<v_float32>::vlanes();
const int stride2 = vlanes_f32 * 2;
v_float64 vs0 = vx_setzero_f64();
v_float64 vs1 = vx_setzero_f64();
v_float64 vs2 = vx_setzero_f64();
v_float64 vs3 = vx_setzero_f64();
for( ; k <= n - stride2; k += stride2 )
{
v_float32 va0 = vx_load(a + k);
v_float32 vb0 = vx_load(b + k);
vs0 = v_muladd(v_cvt_f64(va0), v_cvt_f64(vb0), vs0);
vs1 = v_muladd(v_cvt_f64_high(va0), v_cvt_f64_high(vb0), vs1);
v_float32 va1 = vx_load(a + k + vlanes_f32);
v_float32 vb1 = vx_load(b + k + vlanes_f32);
vs2 = v_muladd(v_cvt_f64(va1), v_cvt_f64(vb1), vs2);
vs3 = v_muladd(v_cvt_f64_high(va1), v_cvt_f64_high(vb1), vs3);
}
vs0 = v_add(v_add(vs0, vs1), v_add(vs2, vs3));
for( ; k <= n - vlanes_f32; k += vlanes_f32 )
{
v_float32 va = vx_load(a + k);
v_float32 vb = vx_load(b + k);
vs0 = v_muladd(v_cvt_f64(va), v_cvt_f64(vb), vs0);
vs0 = v_muladd(v_cvt_f64_high(va), v_cvt_f64_high(vb), vs0);
}
return v_reduce_sum(vs0);
}
// Float input, double accumulation GEMM k-outer/j-inner (for 32FC1 where WT=double)
static inline void simdGEMM_kj_f32f64(
const float* a_data, const float* _b_data, size_t b_step,
const float* _c_data, size_t c_step1,
float* d_data,
int n, int m,
double alpha, double beta,
int& j)
{
const int vlanes_f32 = VTraits<v_float32>::vlanes();
const int vlanes_f64 = VTraits<v_float64>::vlanes();
if( m >= 2 * vlanes_f32 && n >= 64 )
{
// k-outer / j-inner with double accumulation
const int max_m = 1600 / (int)sizeof(float);
cv::AutoBuffer<double> _s_buf(max_m);
double* s_buf = _s_buf.data();
{
v_float64 vz = vx_setzero_f64();
for( j = 0; j <= m - vlanes_f64; j += vlanes_f64 )
v_store(s_buf + j, vz);
for( ; j < m; j++ )
s_buf[j] = 0.0;
}
{
const float* b_ptr = _b_data;
for( int k = 0; k < n; k++, b_ptr += b_step )
{
v_float64 va = vx_setall_f64((double)a_data[k]);
j = 0;
for( ; j <= m - vlanes_f32; j += vlanes_f32 )
{
v_float32 vb = vx_load(b_ptr + j);
v_float64 vb_lo = v_cvt_f64(vb);
v_float64 vb_hi = v_cvt_f64_high(vb);
v_float64 vd_lo = vx_load(s_buf + j);
v_float64 vd_hi = vx_load(s_buf + j + vlanes_f64);
vd_lo = v_muladd(va, vb_lo, vd_lo);
vd_hi = v_muladd(va, vb_hi, vd_hi);
v_store(s_buf + j, vd_lo);
v_store(s_buf + j + vlanes_f64, vd_hi);
}
double a_val = (double)a_data[k];
for( ; j < m; j++ )
s_buf[j] += a_val * (double)b_ptr[j];
}
}
{
v_float64 valpha = vx_setall_f64(alpha);
const float* c_data = _c_data;
j = 0;
if( !_c_data )
{
for( ; j <= m - vlanes_f32; j += vlanes_f32 )
{
v_float64 vs_lo = v_mul(vx_load(s_buf + j), valpha);
v_float64 vs_hi = v_mul(vx_load(s_buf + j + vlanes_f64), valpha);
v_store(d_data + j, v_cvt_f32(vs_lo, vs_hi));
}
for( ; j < m; j++ )
d_data[j] = (float)(s_buf[j] * alpha);
}
else if( c_step1 == 1 )
{
v_float64 vbeta = vx_setall_f64(beta);
for( ; j <= m - vlanes_f32; j += vlanes_f32 )
{
v_float64 vs_lo = v_mul(vx_load(s_buf + j), valpha);
v_float64 vs_hi = v_mul(vx_load(s_buf + j + vlanes_f64), valpha);
v_float32 vc = vx_load(c_data + j);
vs_lo = v_muladd(v_cvt_f64(vc), vbeta, vs_lo);
vs_hi = v_muladd(v_cvt_f64_high(vc), vbeta, vs_hi);
v_store(d_data + j, v_cvt_f32(vs_lo, vs_hi));
}
for( ; j < m; j++ )
d_data[j] = (float)(s_buf[j] * alpha + (double)c_data[j] * beta);
}
else
{
for( j = 0; j < m; j++, c_data += c_step1 )
d_data[j] = (float)(s_buf[j] * alpha + (double)c_data[0] * beta);
}
}
j = m;
}
else if( vlanes_f32 * (int)sizeof(float) >= 32 )
{
// j-outer with SIMD: double accumulation per vlanes_f64 group
const float* c_data = _c_data;
for( ; j <= m - vlanes_f32; j += vlanes_f32, c_data += vlanes_f32 * c_step1 )
{
const float* b = _b_data + j;
v_float64 vs_lo = vx_setzero_f64();
v_float64 vs_hi = vx_setzero_f64();
for( int k = 0; k < n; k++, b += b_step )
{
v_float64 va = vx_setall_f64((double)a_data[k]);
v_float32 vb = vx_load(b);
vs_lo = v_muladd(va, v_cvt_f64(vb), vs_lo);
vs_hi = v_muladd(va, v_cvt_f64_high(vb), vs_hi);
}
v_float64 valpha = vx_setall_f64(alpha);
vs_lo = v_mul(vs_lo, valpha);
vs_hi = v_mul(vs_hi, valpha);
if( !_c_data )
{
v_store(d_data + j, v_cvt_f32(vs_lo, vs_hi));
}
else if( c_step1 == 1 )
{
v_float64 vbeta = vx_setall_f64(beta);
v_float32 vc = vx_load(c_data);
vs_lo = v_muladd(v_cvt_f64(vc), vbeta, vs_lo);
vs_hi = v_muladd(v_cvt_f64_high(vc), vbeta, vs_hi);
v_store(d_data + j, v_cvt_f32(vs_lo, vs_hi));
}
else
{
CV_DECL_ALIGNED(CV_SIMD_WIDTH) float buf[VTraits<v_float32>::max_nlanes];
v_store_aligned(buf, v_cvt_f32(vs_lo, vs_hi));
for( int jj = 0; jj < vlanes_f32; jj++ )
d_data[j + jj] = buf[jj] + c_data[jj * c_step1] * (float)beta;
}
}
}
}
#endif
#if CV_SIMD
template<typename ST>
static inline void simdGEMM_kj(
const ST* a_data, const ST* _b_data, size_t b_step,
const ST* _c_data, size_t c_step1,
ST* d_data,
int n, int m,
double alpha, double beta,
int& j)
{
typedef typename VecTraits<ST>::VT VecT;
const int vlanes = VTraits<VecT>::vlanes();
if( m >= 2 * vlanes && n >= 64 )
{
// k-outer / j-inner SIMD: sequential B-row access, good for large n
const int max_m = 1600 / (int)sizeof(ST);
CV_DECL_ALIGNED(CV_SIMD_WIDTH) ST s_buf[max_m];
{
VecT vz = VecTraits<ST>::zero();
for( j = 0; j <= m - vlanes; j += vlanes )
v_store_aligned(s_buf + j, vz);
for( ; j < m; j++ )
s_buf[j] = (ST)0;
}
{
const ST* b_ptr = _b_data;
for( int k = 0; k < n; k++, b_ptr += b_step )
{
VecT va = VecTraits<ST>::setall(a_data[k]);
j = 0;
for( ; j <= m - vlanes; j += vlanes )
{
VecT vd = vx_load_aligned(s_buf + j);
VecT vb = VecTraits<ST>::load(b_ptr + j);
vd = v_muladd(va, vb, vd);
v_store_aligned(s_buf + j, vd);
}
ST a_val = a_data[k];
for( ; j < m; j++ )
s_buf[j] += a_val * b_ptr[j];
}
}
{
VecT valpha = VecTraits<ST>::setall((ST)alpha);
const ST* c_data = _c_data;
j = 0;
if( !_c_data )
{
for( ; j <= m - vlanes; j += vlanes )
{
VecT vs = v_mul(vx_load_aligned(s_buf + j), valpha);
v_store(d_data + j, vs);
}
for( ; j < m; j++ )
d_data[j] = s_buf[j] * (ST)alpha;
}
else if( c_step1 == 1 )
{
VecT vbeta = VecTraits<ST>::setall((ST)beta);
for( ; j <= m - vlanes; j += vlanes )
{
VecT vs = v_mul(vx_load_aligned(s_buf + j), valpha);
VecT vc = VecTraits<ST>::load(c_data + j);
vs = v_muladd(vc, vbeta, vs);
v_store(d_data + j, vs);
}
for( ; j < m; j++ )
d_data[j] = s_buf[j] * (ST)alpha + c_data[j] * (ST)beta;
}
else
{
for( j = 0; j < m; j++, c_data += c_step1 )
d_data[j] = s_buf[j] * (ST)alpha + c_data[0] * (ST)beta;
}
}
j = m; // all columns handled
}
else if( vlanes * (int)sizeof(ST) >= 32 )
{
// j-outer with SIMD: process vlanes columns per group
// (skip on narrow SIMD where 4-column scalar is faster)
VecT valpha = VecTraits<ST>::setall((ST)alpha);
const ST* c_data = _c_data;
for( ; j <= m - vlanes; j += vlanes, c_data += vlanes * c_step1 )
{
const ST* b = _b_data + j;
VecT vs = VecTraits<ST>::zero();
for( int k = 0; k < n; k++, b += b_step )
{
VecT va = VecTraits<ST>::setall(a_data[k]);
VecT vb = VecTraits<ST>::load(b);
vs = v_muladd(va, vb, vs);
}
vs = v_mul(vs, valpha);
if( !_c_data )
{
v_store(d_data + j, vs);
}
else if( c_step1 == 1 )
{
VecT vc = VecTraits<ST>::load(c_data);
vs = v_muladd(vc, VecTraits<ST>::setall((ST)beta), vs);
v_store(d_data + j, vs);
}
else
{
CV_DECL_ALIGNED(CV_SIMD_WIDTH) ST buf[VTraits<VecT>::max_nlanes];
v_store_aligned(buf, vs);
for( int jj = 0; jj < vlanes; jj++ )
d_data[j + jj] = buf[jj] + c_data[jj * c_step1] * (ST)beta;
}
}
// Remaining columns handled by 4-column scalar in caller
}
}
#endif
template<typename T, typename WT> static void
GEMMSingleMul( const T* a_data, size_t a_step,
@@ -286,20 +618,38 @@ GEMMSingleMul( const T* a_data, size_t a_step,
for( j = 0; j < d_size.width; j++, b_data += b_step,
c_data += c_step1 )
{
WT s0(0), s1(0), s2(0), s3(0);
WT s0(0);
k = 0;
#if CV_ENABLE_UNROLLED
for( ; k <= n - 4; k += 4 )
#if CV_SIMD
if( sizeof(WT) == sizeof(double) )
{
s0 += WT(a_data[k])*WT(b_data[k]);
s1 += WT(a_data[k+1])*WT(b_data[k+1]);
s2 += WT(a_data[k+2])*WT(b_data[k+2]);
s3 += WT(a_data[k+3])*WT(b_data[k+3]);
#if CV_SIMD_64F
if( sizeof(T) == sizeof(double) )
s0 = (WT)simdDotProduct((const double*)a_data, (const double*)b_data, n, k);
else if( sizeof(T) == sizeof(float) )
s0 = (WT)simdDotProduct_f32f64((const float*)a_data, (const float*)b_data, n, k);
else
#endif
s0 = (WT)simdDotProduct((const float*)a_data, (const float*)b_data, n, k);
}
else
#endif
{
WT s1(0), s2(0), s3(0);
#if CV_ENABLE_UNROLLED
for( ; k <= n - 4; k += 4 )
{
s0 += WT(a_data[k])*WT(b_data[k]);
s1 += WT(a_data[k+1])*WT(b_data[k+1]);
s2 += WT(a_data[k+2])*WT(b_data[k+2]);
s3 += WT(a_data[k+3])*WT(b_data[k+3]);
}
#endif
s0 += s1 + s2 + s3;
}
#endif
for( ; k < n; k++ )
s0 += WT(a_data[k])*WT(b_data[k]);
s0 = (s0+s1+s2+s3)*alpha;
s0 = s0*alpha;
if( !c_data )
d_data[j] = T(s0);
@@ -323,7 +673,29 @@ GEMMSingleMul( const T* a_data, size_t a_step,
a_data = a_buf;
}
for( j = 0; j <= m - 4; j += 4, c_data += 4*c_step1 )
j = 0;
#if CV_SIMD
if( sizeof(WT) == sizeof(double) )
{
#if CV_SIMD_64F
if( sizeof(T) == sizeof(double) )
simdGEMM_kj((const double*)a_data, (const double*)_b_data, b_step,
(const double*)_c_data, c_step1,
(double*)d_data, n, m, alpha, beta, j);
else if( sizeof(T) == sizeof(float) )
simdGEMM_kj_f32f64((const float*)a_data, (const float*)_b_data, b_step,
(const float*)_c_data, c_step1,
(float*)d_data, n, m, alpha, beta, j);
else
#endif
simdGEMM_kj((const float*)a_data, (const float*)_b_data, b_step,
(const float*)_c_data, c_step1,
(float*)d_data, n, m, alpha, beta, j);
}
#endif
c_data = _c_data + j * c_step1;
// 4-column j-outer with register accumulators for remaining columns
for( ; j <= m - 4; j += 4, c_data += 4*c_step1 )
{
const T* b = _b_data + j;
WT s0(0), s1(0), s2(0), s3(0);
@@ -394,6 +766,41 @@ GEMMSingleMul( const T* a_data, size_t a_step,
{
WT al(a_data[k]);
j=0;
#if CV_SIMD_64F
if( sizeof(WT) == sizeof(double) )
{
if( sizeof(T) == sizeof(double) )
{
const int vlanes = VTraits<v_float64>::vlanes();
v_float64 val = vx_setall_f64(*((const double*)&al));
for( ; j <= m - vlanes; j += vlanes )
{
v_float64 vd = vx_load((const double*)(d_buf + j));
v_float64 vb = vx_load((const double*)(b_data + j));
vd = v_muladd(vb, val, vd);
v_store((double*)(d_buf + j), vd);
}
}
else
{
const int vlanes_f32 = VTraits<v_float32>::vlanes();
const int vlanes_f64 = VTraits<v_float64>::vlanes();
v_float64 val = vx_setall_f64(*((const double*)&al));
for( ; j <= m - vlanes_f32; j += vlanes_f32 )
{
v_float32 vb = vx_load((const float*)(b_data + j));
v_float64 vb_lo = v_cvt_f64(vb);
v_float64 vb_hi = v_cvt_f64_high(vb);
v_float64 vd_lo = vx_load((const double*)(d_buf + j));
v_float64 vd_hi = vx_load((const double*)(d_buf + j + vlanes_f64));
vd_lo = v_muladd(vb_lo, val, vd_lo);
vd_hi = v_muladd(vb_hi, val, vd_hi);
v_store((double*)(d_buf + j), vd_lo);
v_store((double*)(d_buf + j + vlanes_f64), vd_hi);
}
}
}
#endif
for( ; j < m; j++ )
d_buf[j] += WT(b_data[j])*al;
}
@@ -410,7 +817,96 @@ GEMMSingleMul( const T* a_data, size_t a_step,
}
}
}
#if CV_SIMD_64F
template<typename WT>
static inline void setDoubleDataZero(WT* d_data, int m, int vlanes)
{
v_float64 vz = vx_setzero_f64();
int j = 0;
for( ; j <= m - vlanes; j += vlanes )
v_store((double*)(d_data + j), vz);
for( ; j < m; j++ )
d_data[j] = WT(0);
}
template<typename ST>
static inline void scalarTail_kj(const ST* a_k, const ST* b_j,
double* d_j, int j, int m)
{
double al = (double)(*a_k);
for( ; j < m; j++ )
d_j[j] += al * (double)(b_j[j]);
}
template<typename ST>
static inline void simdMulAdd(const ST* b, double* d, int m, double a_val);
template<>
inline void simdMulAdd<float>(
const float* b, double* d, int m, double a_val)
{
const int vlanes_f32 = VTraits<v_float32>::vlanes();
const int vlanes_f64 = VTraits<v_float64>::vlanes();
v_float64 va = vx_setall_f64(a_val);
int j = 0;
for( ; j <= m - vlanes_f32; j += vlanes_f32 )
{
v_float32 vb = vx_load(b + j);
v_float64 vb_lo = v_cvt_f64(vb);
v_float64 vb_hi = v_cvt_f64_high(vb);
v_float64 vd_lo = vx_load(d + j);
v_float64 vd_hi = vx_load(d + j + vlanes_f64);
vd_lo = v_muladd(va, vb_lo, vd_lo);
vd_hi = v_muladd(va, vb_hi, vd_hi);
v_store(d + j, vd_lo);
v_store(d + j + vlanes_f64, vd_hi);
}
for( ; j < m; j++ )
d[j] += a_val * (double)b[j];
}
template<>
inline void simdMulAdd<double>(
const double* b, double* d, int m, double a_val)
{
const int vlanes = VTraits<v_float64>::vlanes();
const int stride2 = vlanes * 2;
v_float64 va = vx_setall_f64(a_val);
int j = 0;
for( ; j <= m - stride2; j += stride2 )
{
v_float64 vd0 = vx_load(d + j);
v_float64 vb0 = vx_load(b + j);
v_float64 vd1 = vx_load(d + j + vlanes);
v_float64 vb1 = vx_load(b + j + vlanes);
vd0 = v_muladd(va, vb0, vd0);
vd1 = v_muladd(va, vb1, vd1);
v_store(d + j, vd0);
v_store(d + j + vlanes, vd1);
}
for( ; j <= m - vlanes; j += vlanes )
{
v_float64 vd = vx_load(d + j);
v_float64 vb = vx_load(b + j);
vd = v_muladd(va, vb, vd);
v_store(d + j, vd);
}
for( ; j < m; j++ )
d[j] += a_val * b[j];
}
template<typename ST>
static inline void simdBlockMul_kj(
const ST* a_data, const ST* b_data, size_t b_step,
double* d_data, int n, int m, bool do_acc)
{
const int vlanes = VTraits<v_float64>::vlanes();
if( !do_acc )
setDoubleDataZero(d_data, m, vlanes);
for( int k = 0; k < n; k++, b_data += b_step )
simdMulAdd(b_data, d_data, m, (double)a_data[k]);
}
#endif // CV_SIMD_64F
template<typename T, typename WT> static void
GEMMBlockMul( const T* a_data, size_t a_step,
@@ -456,17 +952,35 @@ GEMMBlockMul( const T* a_data, size_t a_step,
for( j = 0; j < d_size.width; j++, b_data += b_step )
{
WT s0 = do_acc ? d_data[j]:WT(0), s1(0);
for( k = 0; k <= n - 2; k += 2 )
WT s0 = do_acc ? d_data[j] : WT(0);
k = 0;
#if CV_SIMD
if( sizeof(WT) == sizeof(double) )
{
s0 += WT(a_data[k])*WT(b_data[k]);
s1 += WT(a_data[k+1])*WT(b_data[k+1]);
#if CV_SIMD_64F
if( sizeof(T) == sizeof(double) )
s0 += (WT)simdDotProduct((const double*)a_data, (const double*)b_data, n, k);
else if( sizeof(T) == sizeof(float) )
s0 += (WT)simdDotProduct_f32f64((const float*)a_data, (const float*)b_data, n, k);
else
#endif
s0 += (WT)simdDotProduct((const float*)a_data, (const float*)b_data, n, k);
}
else
#endif
{
WT s1(0);
for( ; k <= n - 2; k += 2 )
{
s0 += WT(a_data[k])*WT(b_data[k]);
s1 += WT(a_data[k+1])*WT(b_data[k+1]);
}
s0 += s1;
}
for( ; k < n; k++ )
s0 += WT(a_data[k])*WT(b_data[k]);
d_data[j] = s0 + s1;
d_data[j] = s0;
}
}
}
@@ -482,40 +996,55 @@ GEMMBlockMul( const T* a_data, size_t a_step,
a_buf[k] = a_data[a_step1*k];
a_data = a_buf;
}
for( j = 0; j <= m - 4; j += 4 )
#if CV_SIMD_64F
if( sizeof(WT) == sizeof(double) )
{
WT s0, s1, s2, s3;
const T* b = b_data + j;
if( do_acc )
{
s0 = d_data[j]; s1 = d_data[j+1];
s2 = d_data[j+2]; s3 = d_data[j+3];
}
if( sizeof(T) == sizeof(double) )
simdBlockMul_kj(
(const double*)a_data, (const double*)b_data, b_step,
(double*)d_data, n, m, do_acc != 0);
else
s0 = s1 = s2 = s3 = WT(0);
for( k = 0; k < n; k++, b += b_step )
simdBlockMul_kj(
(const float*)a_data, (const float*)b_data, b_step,
(double*)d_data, n, m, do_acc != 0);
}
else
#endif
{
for( j = 0; j <= m - 4; j += 4 )
{
WT a(a_data[k]);
s0 += a * WT(b[0]); s1 += a * WT(b[1]);
s2 += a * WT(b[2]); s3 += a * WT(b[3]);
WT s0, s1, s2, s3;
const T* b = b_data + j;
if( do_acc )
{
s0 = d_data[j]; s1 = d_data[j+1];
s2 = d_data[j+2]; s3 = d_data[j+3];
}
else
s0 = s1 = s2 = s3 = WT(0);
for( k = 0; k < n; k++, b += b_step )
{
WT a(a_data[k]);
s0 += a * WT(b[0]); s1 += a * WT(b[1]);
s2 += a * WT(b[2]); s3 += a * WT(b[3]);
}
d_data[j] = s0; d_data[j+1] = s1;
d_data[j+2] = s2; d_data[j+3] = s3;
}
d_data[j] = s0; d_data[j+1] = s1;
d_data[j+2] = s2; d_data[j+3] = s3;
}
for( ; j < m; j++ )
{
const T* b = b_data + j;
WT s0 = do_acc ? d_data[j] : WT(0);
for( ; j < m; j++ )
{
const T* b = b_data + j;
WT s0 = do_acc ? d_data[j] : WT(0);
for( k = 0; k < n; k++, b += b_step )
s0 += WT(a_data[k]) * WT(b[0]);
for( k = 0; k < n; k++, b += b_step )
s0 += WT(a_data[k]) * WT(b[0]);
d_data[j] = s0;
d_data[j] = s0;
}
}
}
}
@@ -680,7 +1209,7 @@ static void GEMMBlockMul_32fc( const Complexf* a_data, size_t a_step,
Complexd* d_data, size_t d_step,
Size a_size, Size d_size, int flags )
{
GEMMBlockMul(a_data, a_step, b_data, b_step, d_data, d_step, a_size, d_size, flags);
GEMMBlockMul<Complexf,Complexd>(a_data, a_step, b_data, b_step, d_data, d_step, a_size, d_size, flags);
}
@@ -689,7 +1218,7 @@ static void GEMMBlockMul_64fc( const Complexd* a_data, size_t a_step,
Complexd* d_data, size_t d_step,
Size a_size, Size d_size, int flags )
{
GEMMBlockMul(a_data, a_step, b_data, b_step, d_data, d_step, a_size, d_size, flags);
GEMMBlockMul<Complexd,Complexd>(a_data, a_step, b_data, b_step, d_data, d_step, a_size, d_size, flags);
}