Extract AVX code from engine_util_blas.c into engine_util_blas_avx.h
PiperOrigin-RevId: 911274884 Change-Id: Iaac2789f9538f669c49b497d1e60947a783368d3
This commit is contained in:
committed by
Copybara-Service
parent
579a27e9d2
commit
6376e67070
+18
-257
@@ -18,12 +18,7 @@
|
||||
|
||||
#include <mujoco/mjtnum.h>
|
||||
|
||||
#ifdef mjUSEPLATFORMSIMD
|
||||
#if defined(__AVX__) && !defined(mjUSESINGLE)
|
||||
#define mjUSEAVX
|
||||
#include "immintrin.h"
|
||||
#endif
|
||||
#endif
|
||||
#include "engine/engine_util_blas_avx.h" // IWYU pragma: keep
|
||||
|
||||
|
||||
|
||||
@@ -339,42 +334,12 @@ void mju_scl(mjtNum* res, const mjtNum* vec, mjtNum scl, int n) {
|
||||
int i = 0;
|
||||
|
||||
#ifdef mjUSEAVX
|
||||
int n_4 = n - 4;
|
||||
i = mju_scl_avx(res, vec, scl, n);
|
||||
#endif
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d sclpar, val1, val1scl;
|
||||
|
||||
// init
|
||||
sclpar = _mm256_set1_pd(scl);
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(vec+i);
|
||||
val1scl = _mm256_mul_pd(val1, sclpar);
|
||||
_mm256_storeu_pd(res+i, val1scl);
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
|
||||
// process remaining
|
||||
int n_i = n - i;
|
||||
if (n_i == 3) {
|
||||
res[i] = vec[i]*scl;
|
||||
res[i+1] = vec[i+1]*scl;
|
||||
res[i+2] = vec[i+2]*scl;
|
||||
} else if (n_i == 2) {
|
||||
res[i] = vec[i]*scl;
|
||||
res[i+1] = vec[i+1]*scl;
|
||||
} else if (n_i == 1) {
|
||||
res[i] = vec[i]*scl;
|
||||
}
|
||||
|
||||
#else
|
||||
for (; i < n; i++) {
|
||||
res[i] = vec[i]*scl;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -383,40 +348,12 @@ void mju_add(mjtNum* res, const mjtNum* vec1, const mjtNum* vec2, int n) {
|
||||
int i = 0;
|
||||
|
||||
#ifdef mjUSEAVX
|
||||
int n_4 = n - 4;
|
||||
i = mju_add_avx(res, vec1, vec2, n);
|
||||
#endif
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d sum, val1, val2;
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(vec1+i);
|
||||
val2 = _mm256_loadu_pd(vec2+i);
|
||||
sum = _mm256_add_pd(val1, val2);
|
||||
_mm256_storeu_pd(res+i, sum);
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
|
||||
// process remaining
|
||||
int n_i = n - i;
|
||||
if (n_i == 3) {
|
||||
res[i] = vec1[i] + vec2[i];
|
||||
res[i+1] = vec1[i+1] + vec2[i+1];
|
||||
res[i+2] = vec1[i+2] + vec2[i+2];
|
||||
} else if (n_i == 2) {
|
||||
res[i] = vec1[i] + vec2[i];
|
||||
res[i+1] = vec1[i+1] + vec2[i+1];
|
||||
} else if (n_i == 1) {
|
||||
res[i] = vec1[i] + vec2[i];
|
||||
}
|
||||
|
||||
#else
|
||||
for (; i < n; i++) {
|
||||
res[i] = vec1[i] + vec2[i];
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -434,40 +371,12 @@ void mju_sub(mjtNum* res, const mjtNum* vec1, const mjtNum* vec2, int n) {
|
||||
int i = 0;
|
||||
|
||||
#ifdef mjUSEAVX
|
||||
int n_4 = n - 4;
|
||||
i = mju_sub_avx(res, vec1, vec2, n);
|
||||
#endif
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d dif, val1, val2;
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(vec1+i);
|
||||
val2 = _mm256_loadu_pd(vec2+i);
|
||||
dif = _mm256_sub_pd(val1, val2);
|
||||
_mm256_storeu_pd(res+i, dif);
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
|
||||
// process remaining
|
||||
int n_i = n - i;
|
||||
if (n_i == 3) {
|
||||
res[i] = vec1[i] - vec2[i];
|
||||
res[i+1] = vec1[i+1] - vec2[i+1];
|
||||
res[i+2] = vec1[i+2] - vec2[i+2];
|
||||
} else if (n_i == 2) {
|
||||
res[i] = vec1[i] - vec2[i];
|
||||
res[i+1] = vec1[i+1] - vec2[i+1];
|
||||
} else if (n_i == 1) {
|
||||
res[i] = vec1[i] - vec2[i];
|
||||
}
|
||||
|
||||
#else
|
||||
for (; i < n; i++) {
|
||||
res[i] = vec1[i] - vec2[i];
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -485,40 +394,12 @@ void mju_addTo(mjtNum* res, const mjtNum* vec, int n) {
|
||||
int i = 0;
|
||||
|
||||
#ifdef mjUSEAVX
|
||||
int n_4 = n - 4;
|
||||
i = mju_addTo_avx(res, vec, n);
|
||||
#endif
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d sum, val1, val2;
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(res+i);
|
||||
val2 = _mm256_loadu_pd(vec+i);
|
||||
sum = _mm256_add_pd(val1, val2);
|
||||
_mm256_storeu_pd(res+i, sum);
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
|
||||
// process remaining
|
||||
int n_i = n - i;
|
||||
if (n_i == 3) {
|
||||
res[i] += vec[i];
|
||||
res[i+1] += vec[i+1];
|
||||
res[i+2] += vec[i+2];
|
||||
} else if (n_i == 2) {
|
||||
res[i] += vec[i];
|
||||
res[i+1] += vec[i+1];
|
||||
} else if (n_i == 1) {
|
||||
res[i] += vec[i];
|
||||
}
|
||||
|
||||
#else
|
||||
for (; i < n; i++) {
|
||||
res[i] += vec[i];
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -536,40 +417,12 @@ void mju_subFrom(mjtNum* res, const mjtNum* vec, int n) {
|
||||
int i = 0;
|
||||
|
||||
#ifdef mjUSEAVX
|
||||
int n_4 = n - 4;
|
||||
i = mju_subFrom_avx(res, vec, n);
|
||||
#endif
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d dif, val1, val2;
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(res+i);
|
||||
val2 = _mm256_loadu_pd(vec+i);
|
||||
dif = _mm256_sub_pd(val1, val2);
|
||||
_mm256_storeu_pd(res+i, dif);
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
|
||||
// process remaining
|
||||
int n_i = n - i;
|
||||
if (n_i == 3) {
|
||||
res[i] -= vec[i];
|
||||
res[i+1] -= vec[i+1];
|
||||
res[i+2] -= vec[i+2];
|
||||
} else if (n_i == 2) {
|
||||
res[i] -= vec[i];
|
||||
res[i+1] -= vec[i+1];
|
||||
} else if (n_i == 1) {
|
||||
res[i] -= vec[i];
|
||||
}
|
||||
|
||||
#else
|
||||
for (; i < n; i++) {
|
||||
res[i] -= vec[i];
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -578,44 +431,12 @@ void mju_addToScl(mjtNum* res, const mjtNum* vec, mjtNum scl, int n) {
|
||||
int i = 0;
|
||||
|
||||
#ifdef mjUSEAVX
|
||||
int n_4 = n - 4;
|
||||
i = mju_addToScl_avx(res, vec, scl, n);
|
||||
#endif
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d sclpar, sum, val1, val2, val2scl;
|
||||
|
||||
// init
|
||||
sclpar = _mm256_set1_pd(scl);
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(res+i);
|
||||
val2 = _mm256_loadu_pd(vec+i);
|
||||
val2scl = _mm256_mul_pd(val2, sclpar);
|
||||
sum = _mm256_add_pd(val1, val2scl);
|
||||
_mm256_storeu_pd(res+i, sum);
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
|
||||
// process remaining
|
||||
int n_i = n - i;
|
||||
if (n_i == 3) {
|
||||
res[i] += vec[i]*scl;
|
||||
res[i+1] += vec[i+1]*scl;
|
||||
res[i+2] += vec[i+2]*scl;
|
||||
} else if (n_i == 2) {
|
||||
res[i] += vec[i]*scl;
|
||||
res[i+1] += vec[i+1]*scl;
|
||||
} else if (n_i == 1) {
|
||||
res[i] += vec[i]*scl;
|
||||
}
|
||||
|
||||
#else
|
||||
for (; i < n; i++) {
|
||||
res[i] += vec[i]*scl;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -634,45 +455,13 @@ void mju_addToSclInd(mjtNum* res, const mjtNum* vec, const int* ind, mjtNum scl,
|
||||
void mju_addScl(mjtNum* res, const mjtNum* vec1, const mjtNum* vec2, mjtNum scl, int n) {
|
||||
int i = 0;
|
||||
|
||||
#if defined(__AVX__) && defined(mjUSEAVX) && !defined(mjUSESINGLE)
|
||||
int n_4 = n - 4;
|
||||
#ifdef mjUSEAVX
|
||||
i = mju_addScl_avx(res, vec1, vec2, scl, n);
|
||||
#endif
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d sclpar, sum, val1, val2, val2scl;
|
||||
|
||||
// init
|
||||
sclpar = _mm256_set1_pd(scl);
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(vec1+i);
|
||||
val2 = _mm256_loadu_pd(vec2+i);
|
||||
val2scl = _mm256_mul_pd(val2, sclpar);
|
||||
sum = _mm256_add_pd(val1, val2scl);
|
||||
_mm256_storeu_pd(res+i, sum);
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
|
||||
// process remaining
|
||||
int n_i = n - i;
|
||||
if (n_i == 3) {
|
||||
res[i] = vec1[i] + vec2[i]*scl;
|
||||
res[i+1] = vec1[i+1] + vec2[i+1]*scl;
|
||||
res[i+2] = vec1[i+2] + vec2[i+2]*scl;
|
||||
} else if (n_i == 2) {
|
||||
res[i] = vec1[i] + vec2[i]*scl;
|
||||
res[i+1] = vec1[i+1] + vec2[i+1]*scl;
|
||||
} else if (n_i == 1) {
|
||||
res[i] = vec1[i] + vec2[i]*scl;
|
||||
}
|
||||
|
||||
#else
|
||||
for (; i < n; i++) {
|
||||
res[i] = vec1[i] + vec2[i]*scl;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -704,41 +493,13 @@ mjtNum mju_norm(const mjtNum* res, int n) {
|
||||
mjtNum mju_dot(const mjtNum* vec1, const mjtNum* vec2, int n) {
|
||||
mjtNum res = 0;
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
#ifdef mjUSEAVX
|
||||
|
||||
// vector part
|
||||
if (n_4 >= 0) {
|
||||
__m256d sum, prod, val1, val2;
|
||||
__m128d vlow, vhigh, high64;
|
||||
|
||||
// init
|
||||
val1 = _mm256_loadu_pd(vec1);
|
||||
val2 = _mm256_loadu_pd(vec2);
|
||||
sum = _mm256_mul_pd(val1, val2);
|
||||
i = 4;
|
||||
|
||||
// parallel computation
|
||||
while (i <= n_4) {
|
||||
val1 = _mm256_loadu_pd(vec1+i);
|
||||
val2 = _mm256_loadu_pd(vec2+i);
|
||||
prod = _mm256_mul_pd(val1, val2);
|
||||
sum = _mm256_add_pd(sum, prod);
|
||||
i += 4;
|
||||
}
|
||||
|
||||
// reduce
|
||||
vlow = _mm256_castpd256_pd128(sum);
|
||||
vhigh = _mm256_extractf128_pd(sum, 1);
|
||||
vlow = _mm_add_pd(vlow, vhigh);
|
||||
high64 = _mm_unpackhi_pd(vlow, vlow);
|
||||
res = _mm_cvtsd_f64(_mm_add_sd(vlow, high64));
|
||||
}
|
||||
|
||||
res = mju_dot_avx(vec1, vec2, n, &i);
|
||||
#else
|
||||
// do the same order of additions as the AVX intrinsics implementation.
|
||||
// this is faster than the simple for loop you'd expect for a dot product,
|
||||
// and produces exactly the same results.
|
||||
int n_4 = n - 4;
|
||||
mjtNum res0 = 0;
|
||||
mjtNum res1 = 0;
|
||||
mjtNum res2 = 0;
|
||||
|
||||
Reference in New Issue
Block a user