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
@@ -77,6 +77,7 @@ set(MUJOCO_ENGINE_SRCS
|
||||
engine_support.h
|
||||
engine_util_blas.c
|
||||
engine_util_blas.h
|
||||
engine_util_blas_avx.h
|
||||
engine_util_errmem.c
|
||||
engine_util_errmem.h
|
||||
engine_util_misc.c
|
||||
|
||||
+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;
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
// Copyright 2026 DeepMind Technologies Limited
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#ifndef MUJOCO_SRC_ENGINE_ENGINE_UTIL_BLAS_AVX_H_
|
||||
#define MUJOCO_SRC_ENGINE_ENGINE_UTIL_BLAS_AVX_H_
|
||||
|
||||
#ifdef mjUSEPLATFORMSIMD
|
||||
#if defined(__AVX__) && !defined(mjUSESINGLE)
|
||||
|
||||
#define mjUSEAVX
|
||||
|
||||
#include <immintrin.h>
|
||||
|
||||
#include <mujoco/mjtnum.h>
|
||||
|
||||
// res = vec*scl
|
||||
static inline
|
||||
int mju_scl_avx(mjtNum* res, const mjtNum* vec, mjtNum scl, int n) {
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
__m256d sclpar = _mm256_set1_pd(scl);
|
||||
while (i <= n_4) {
|
||||
_mm256_storeu_pd(res+i, _mm256_mul_pd(_mm256_loadu_pd(vec+i), sclpar));
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
return i;
|
||||
}
|
||||
|
||||
// res = vec1 + vec2
|
||||
static inline
|
||||
int mju_add_avx(mjtNum* res, const mjtNum* vec1, const mjtNum* vec2, int n) {
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
while (i <= n_4) {
|
||||
_mm256_storeu_pd(res+i, _mm256_add_pd(_mm256_loadu_pd(vec1+i), _mm256_loadu_pd(vec2+i)));
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
return i;
|
||||
}
|
||||
|
||||
// res = vec1 - vec2
|
||||
static inline
|
||||
int mju_sub_avx(mjtNum* res, const mjtNum* vec1, const mjtNum* vec2, int n) {
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
while (i <= n_4) {
|
||||
_mm256_storeu_pd(res+i, _mm256_sub_pd(_mm256_loadu_pd(vec1+i), _mm256_loadu_pd(vec2+i)));
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
return i;
|
||||
}
|
||||
|
||||
// res += vec
|
||||
static inline
|
||||
int mju_addTo_avx(mjtNum* res, const mjtNum* vec, int n) {
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
while (i <= n_4) {
|
||||
_mm256_storeu_pd(res+i, _mm256_add_pd(_mm256_loadu_pd(res+i), _mm256_loadu_pd(vec+i)));
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
return i;
|
||||
}
|
||||
|
||||
// res -= vec
|
||||
static inline
|
||||
int mju_subFrom_avx(mjtNum* res, const mjtNum* vec, int n) {
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
while (i <= n_4) {
|
||||
_mm256_storeu_pd(res+i, _mm256_sub_pd(_mm256_loadu_pd(res+i), _mm256_loadu_pd(vec+i)));
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
return i;
|
||||
}
|
||||
|
||||
// res += vec*scl
|
||||
static inline
|
||||
int mju_addToScl_avx(mjtNum* res, const mjtNum* vec, mjtNum scl, int n) {
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
__m256d sclpar = _mm256_set1_pd(scl);
|
||||
while (i <= n_4) {
|
||||
__m256d val1 = _mm256_loadu_pd(res+i);
|
||||
__m256d val2 = _mm256_loadu_pd(vec+i);
|
||||
_mm256_storeu_pd(res+i, _mm256_add_pd(val1, _mm256_mul_pd(val2, sclpar)));
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
return i;
|
||||
}
|
||||
|
||||
// res = vec1 + vec2*scl
|
||||
static inline
|
||||
int mju_addScl_avx(mjtNum* res, const mjtNum* vec1, const mjtNum* vec2, mjtNum scl, int n) {
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
__m256d sclpar = _mm256_set1_pd(scl);
|
||||
while (i <= n_4) {
|
||||
__m256d val1 = _mm256_loadu_pd(vec1+i);
|
||||
__m256d val2 = _mm256_loadu_pd(vec2+i);
|
||||
_mm256_storeu_pd(res+i, _mm256_add_pd(val1, _mm256_mul_pd(val2, sclpar)));
|
||||
i += 4;
|
||||
}
|
||||
}
|
||||
return i;
|
||||
}
|
||||
|
||||
// vector dot-product
|
||||
static inline
|
||||
mjtNum mju_dot_avx(const mjtNum* vec1, const mjtNum* vec2, int n, int* processed) {
|
||||
mjtNum res = 0;
|
||||
int i = 0;
|
||||
int n_4 = n - 4;
|
||||
if (n_4 >= 0) {
|
||||
__m256d sum = _mm256_mul_pd(_mm256_loadu_pd(vec1), _mm256_loadu_pd(vec2));
|
||||
i = 4;
|
||||
|
||||
while (i <= n_4) {
|
||||
sum = _mm256_add_pd(sum, _mm256_mul_pd(_mm256_loadu_pd(vec1+i), _mm256_loadu_pd(vec2+i)));
|
||||
i += 4;
|
||||
}
|
||||
|
||||
__m128d vlow = _mm256_castpd256_pd128(sum);
|
||||
__m128d vhigh = _mm256_extractf128_pd(sum, 1);
|
||||
vlow = _mm_add_pd(vlow, vhigh);
|
||||
__m128d high64 = _mm_unpackhi_pd(vlow, vlow);
|
||||
res = _mm_cvtsd_f64(_mm_add_sd(vlow, high64));
|
||||
}
|
||||
*processed = i;
|
||||
return res;
|
||||
}
|
||||
|
||||
#endif // defined(__AVX__) && !defined(mjUSESINGLE)
|
||||
#endif // mjUSEPLATFORMSIMD
|
||||
|
||||
#endif // MUJOCO_SRC_ENGINE_ENGINE_UTIL_BLAS_AVX_H_
|
||||
Reference in New Issue
Block a user