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:
Yuval Tassa
2026-05-06 05:25:15 -07:00
committed by Copybara-Service
parent 579a27e9d2
commit 6376e67070
3 changed files with 179 additions and 257 deletions
+1
View File
@@ -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
View File
@@ -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;
+160
View File
@@ -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_