From 6376e6707036790f09b5f6e20e8dc5acd98d3542 Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Wed, 6 May 2026 05:25:15 -0700 Subject: [PATCH] Extract AVX code from engine_util_blas.c into engine_util_blas_avx.h PiperOrigin-RevId: 911274884 Change-Id: Iaac2789f9538f669c49b497d1e60947a783368d3 --- src/engine/CMakeLists.txt | 1 + src/engine/engine_util_blas.c | 275 ++---------------------------- src/engine/engine_util_blas_avx.h | 160 +++++++++++++++++ 3 files changed, 179 insertions(+), 257 deletions(-) create mode 100644 src/engine/engine_util_blas_avx.h diff --git a/src/engine/CMakeLists.txt b/src/engine/CMakeLists.txt index 3758ed70..e61c1dc4 100644 --- a/src/engine/CMakeLists.txt +++ b/src/engine/CMakeLists.txt @@ -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 diff --git a/src/engine/engine_util_blas.c b/src/engine/engine_util_blas.c index 0041484f..8d4e7c3e 100644 --- a/src/engine/engine_util_blas.c +++ b/src/engine/engine_util_blas.c @@ -18,12 +18,7 @@ #include -#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; diff --git a/src/engine/engine_util_blas_avx.h b/src/engine/engine_util_blas_avx.h new file mode 100644 index 00000000..ff2af6c7 --- /dev/null +++ b/src/engine/engine_util_blas_avx.h @@ -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 + +#include + +// 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_