diff --git a/src/engine/CMakeLists.txt b/src/engine/CMakeLists.txt index fe09f920..1c0deffa 100644 --- a/src/engine/CMakeLists.txt +++ b/src/engine/CMakeLists.txt @@ -64,6 +64,7 @@ set(MUJOCO_ENGINE_SRCS engine_util_solve.h engine_util_sparse.c engine_util_sparse.h + engine_util_sparse_avx.h engine_util_spatial.c engine_util_spatial.h engine_vfs.c diff --git a/src/engine/engine_util_sparse.c b/src/engine/engine_util_sparse.c index 6c7f869f..453702a3 100644 --- a/src/engine/engine_util_sparse.c +++ b/src/engine/engine_util_sparse.c @@ -13,6 +13,7 @@ // limitations under the License. #include "engine/engine_util_sparse.h" +#include "engine/engine_util_sparse_avx.h" #include @@ -22,61 +23,18 @@ #include "engine/engine_macro.h" #include "engine/engine_util_blas.h" -#ifdef mjUSEPLATFORMSIMD - #if defined(__AVX__) && defined(mjUSEDOUBLE) - #define mjUSEAVX - #include "immintrin.h" - #endif -#endif //------------------------------ sparse operations ------------------------------------------------- // dot-product, first vector is sparse mjtNum mju_dotSparse(const mjtNum* vec1, const mjtNum* vec2, const int nnz1, const int* ind1) { +#ifdef mjUSEAVX + return mju_dotSparse_avx(vec1, vec2, nnz1, ind1); +#else int i = 0; mjtNum res = 0; - -#ifdef mjUSEAVX - int nnz1_4 = nnz1 - 4; - - // vector part - if (nnz1_4>=0) { - __m256d sum, prod, val1, val2; - __m128d vlow, vhigh, high64; - - // init - val2 = _mm256_set_pd(vec2[ind1[3]], - vec2[ind1[2]], - vec2[ind1[1]], - vec2[ind1[0]]); - val1 = _mm256_loadu_pd(vec1); - sum = _mm256_mul_pd(val1, val2); - i = 4; - - // parallel computation - while (i<=nnz1_4) { - val1 = _mm256_loadu_pd(vec1+i); - val2 = _mm256_set_pd(vec2[ind1[i+3]], - vec2[ind1[i+2]], - vec2[ind1[i+1]], - vec2[ind1[i+0]]); - 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)); - } - -#else int n_4 = nnz1 - 4; - mjtNum res0 = 0; mjtNum res1 = 0; mjtNum res2 = 0; @@ -89,7 +47,6 @@ mjtNum mju_dotSparse(const mjtNum* vec1, const mjtNum* vec2, res3 += vec1[i+3] * vec2[ind1[i+3]]; } res = (res0 + res2) + (res1 + res3); -#endif // scalar part for (; i=0) { - __m256d sum0, sum1, sum2, prod, val1, val2; - __m128d vlow, vhigh, high64; - - // init - val2 = _mm256_set_pd(vec2[ind1[3]], - vec2[ind1[2]], - vec2[ind1[1]], - vec2[ind1[0]]); - val1 = _mm256_loadu_pd(vec10); - sum0 = _mm256_mul_pd(val1, val2); - val1 = _mm256_loadu_pd(vec11); - sum1 = _mm256_mul_pd(val1, val2); - val1 = _mm256_loadu_pd(vec12); - sum2 = _mm256_mul_pd(val1, val2); - i = 4; - - // parallel computation - while (i<=nnz1_4) { - // get val2 only once - val2 = _mm256_set_pd(vec2[ind1[i+3]], - vec2[ind1[i+2]], - vec2[ind1[i+1]], - vec2[ind1[i+0]]); - - // process each val1 - val1 = _mm256_loadu_pd(vec10+i); - prod = _mm256_mul_pd(val1, val2); - sum0 = _mm256_add_pd(sum0, prod); - - val1 = _mm256_loadu_pd(vec11+i); - prod = _mm256_mul_pd(val1, val2); - sum1 = _mm256_add_pd(sum1, prod); - - val1 = _mm256_loadu_pd(vec12+i); - prod = _mm256_mul_pd(val1, val2); - sum2 = _mm256_add_pd(sum2, prod); - - i += 4; - } - - // reduce - vlow = _mm256_castpd256_pd128(sum0); - vhigh = _mm256_extractf128_pd(sum0, 1); - vlow = _mm_add_pd(vlow, vhigh); - high64 = _mm_unpackhi_pd(vlow, vlow); - RES0 = _mm_cvtsd_f64(_mm_add_sd(vlow, high64)); - - vlow = _mm256_castpd256_pd128(sum1); - vhigh = _mm256_extractf128_pd(sum1, 1); - vlow = _mm_add_pd(vlow, vhigh); - high64 = _mm_unpackhi_pd(vlow, vlow); - RES1 = _mm_cvtsd_f64(_mm_add_sd(vlow, high64)); - - vlow = _mm256_castpd256_pd128(sum2); - vhigh = _mm256_extractf128_pd(sum2, 1); - vlow = _mm_add_pd(vlow, vhigh); - high64 = _mm_unpackhi_pd(vlow, vlow); - RES2 = _mm_cvtsd_f64(_mm_add_sd(vlow, high64)); - } -#endif - - // scalar part for (; i=3) { - mju_dotSparseX3(res+r, res+r+1, res+r+2, - mat+rowadr[r], mat+rowadr[r+1], mat+rowadr[r+2], - vec, rownnz[r], colind+rowadr[r]); - - r += 3; - rs -= 3; - } - - // handle remaining rows - while (rs>0) { - res[r] = mju_dotSparse(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]); - - r++; - rs--; - } - - // go back one, because of outer for loop - r--; - } - - else { - res[r] = mju_dotSparse(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]); - } + res[r] = mju_dotSparse(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]); } +#endif // mjUSEAVX } + + // res = res*scl1 + vec*scl2 static void mju_addToSclScl(mjtNum* res, const mjtNum* vec, mjtNum scl1, mjtNum scl2, int n) { - int i = 0; - #ifdef mjUSEAVX - int n_4 = n - 4; - - // vector part - if (n_4>=0) { - __m256d sclpar1, sclpar2, sum, val1, val2; - - // init - sclpar1 = _mm256_set1_pd(scl1); - sclpar2 = _mm256_set1_pd(scl2); - - // parallel computation - while (i<=n_4) { - val1 = _mm256_loadu_pd(res+i); - val2 = _mm256_loadu_pd(vec+i); - val1 = _mm256_mul_pd(val1, sclpar1); - val2 = _mm256_mul_pd(val2, sclpar2); - 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] = res[i]*scl1 + vec[i]*scl2; - res[i+1] = res[i+1]*scl1 + vec[i+1]*scl2; - res[i+2] = res[i+2]*scl1 + vec[i+2]*scl2; - } else if (n_i==2) { - res[i] = res[i]*scl1 + vec[i]*scl2; - res[i+1] = res[i+1]*scl1 + vec[i+1]*scl2; - } else if (n_i==1) { - res[i] = res[i]*scl1 + vec[i]*scl2; - } - + mju_addToSclScl_avx(res, vec, scl1, scl2, n); #else - for (; i=0) { - __m128i val1, val2, cmp; - - // parallel computation - while (i<=n_4) { - val1 = _mm_loadu_si128((const __m128i*)(vec1+i)); - val2 = _mm_loadu_si128((const __m128i*)(vec2+i)); - cmp = _mm_cmpeq_epi32(val1, val2); - if (_mm_movemask_epi8(cmp)!= 0xFFFF) { - return 0; - } - i += 4; - } - } -#endif - - // scalar part - return !memcmp(vec1+i, vec2+i, (n-i)*sizeof(int)); + return mju_compare_avx(vec1, vec2, n); +#else + return !memcmp(vec1, vec2, n*sizeof(int)); +#endif // mjUSEAVX } + // combine two sparse vectors: dst = a*dst + b*src, return nnz of result int mju_combineSparse(mjtNum* dst, const mjtNum* src, int n, mjtNum a, mjtNum b, int dst_nnz, int src_nnz, int* dst_ind, const int* src_ind, diff --git a/src/engine/engine_util_sparse_avx.h b/src/engine/engine_util_sparse_avx.h new file mode 100644 index 00000000..0a8db47d --- /dev/null +++ b/src/engine/engine_util_sparse_avx.h @@ -0,0 +1,292 @@ +// Copyright 2023 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_SPARSE_AVX_H_ +#define MUJOCO_SRC_ENGINE_ENGINE_UTIL_SPARSE_AVX_H_ + +#ifdef mjUSEPLATFORMSIMD +#if defined(__AVX__) && defined(mjUSEDOUBLE) + +#define mjUSEAVX + +#include +#include + +#include +#include "engine/engine_macro.h" + + +//------------------------------ sparse operations using avx --------------------------------------- + +// dot-product, first vector is sparse +static inline +mjtNum mju_dotSparse_avx(const mjtNum* vec1, const mjtNum* vec2, + const int nnz1, const int* ind1) { + int i = 0; + mjtNum res = 0; + int nnz1_4 = nnz1 - 4; + + // vector part + if (nnz1_4>=0) { + __m256d sum, prod, val1, val2; + __m128d vlow, vhigh, high64; + + // init + val2 = _mm256_set_pd(vec2[ind1[3]], + vec2[ind1[2]], + vec2[ind1[1]], + vec2[ind1[0]]); + val1 = _mm256_loadu_pd(vec1); + sum = _mm256_mul_pd(val1, val2); + i = 4; + + // parallel computation + while (i<=nnz1_4) { + val1 = _mm256_loadu_pd(vec1+i); + val2 = _mm256_set_pd(vec2[ind1[i+3]], + vec2[ind1[i+2]], + vec2[ind1[i+1]], + vec2[ind1[i+0]]); + 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)); + } + + // scalar part + for (; i=0) { + __m256d sum0, sum1, sum2, prod, val1, val2; + __m128d vlow, vhigh, high64; + + // init + val2 = _mm256_set_pd(vec2[ind1[3]], + vec2[ind1[2]], + vec2[ind1[1]], + vec2[ind1[0]]); + val1 = _mm256_loadu_pd(vec10); + sum0 = _mm256_mul_pd(val1, val2); + val1 = _mm256_loadu_pd(vec11); + sum1 = _mm256_mul_pd(val1, val2); + val1 = _mm256_loadu_pd(vec12); + sum2 = _mm256_mul_pd(val1, val2); + i = 4; + + // parallel computation + while (i<=nnz1_4) { + // get val2 only once + val2 = _mm256_set_pd(vec2[ind1[i+3]], + vec2[ind1[i+2]], + vec2[ind1[i+1]], + vec2[ind1[i+0]]); + + // process each val1 + val1 = _mm256_loadu_pd(vec10+i); + prod = _mm256_mul_pd(val1, val2); + sum0 = _mm256_add_pd(sum0, prod); + + val1 = _mm256_loadu_pd(vec11+i); + prod = _mm256_mul_pd(val1, val2); + sum1 = _mm256_add_pd(sum1, prod); + + val1 = _mm256_loadu_pd(vec12+i); + prod = _mm256_mul_pd(val1, val2); + sum2 = _mm256_add_pd(sum2, prod); + + i += 4; + } + + // reduce + vlow = _mm256_castpd256_pd128(sum0); + vhigh = _mm256_extractf128_pd(sum0, 1); + vlow = _mm_add_pd(vlow, vhigh); + high64 = _mm_unpackhi_pd(vlow, vlow); + RES0 = _mm_cvtsd_f64(_mm_add_sd(vlow, high64)); + + vlow = _mm256_castpd256_pd128(sum1); + vhigh = _mm256_extractf128_pd(sum1, 1); + vlow = _mm_add_pd(vlow, vhigh); + high64 = _mm_unpackhi_pd(vlow, vlow); + RES1 = _mm_cvtsd_f64(_mm_add_sd(vlow, high64)); + + vlow = _mm256_castpd256_pd128(sum2); + vhigh = _mm256_extractf128_pd(sum2, 1); + vlow = _mm_add_pd(vlow, vhigh); + high64 = _mm_unpackhi_pd(vlow, vlow); + RES2 = _mm_cvtsd_f64(_mm_add_sd(vlow, high64)); + } + + // scalar part + for (; i=3) { + mju_dotSparseX3_avx(res+r, res+r+1, res+r+2, + mat+rowadr[r], mat+rowadr[r+1], mat+rowadr[r+2], + vec, rownnz[r], colind+rowadr[r]); + + r += 3; + rs -= 3; + } + + // handle remaining rows + while (rs>0) { + res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]); + + r++; + rs--; + } + + // go back one, because of outer for loop + r--; + } + + else { + res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]); + } + } +} + +// res = res*scl1 + vec*scl2 +static inline +void mju_addToSclScl_avx(mjtNum* res, const mjtNum* vec, mjtNum scl1, mjtNum scl2, int n) { + int i = 0; + + int n_4 = n - 4; + + // vector part + if (n_4>=0) { + __m256d sclpar1, sclpar2, sum, val1, val2; + + // init + sclpar1 = _mm256_set1_pd(scl1); + sclpar2 = _mm256_set1_pd(scl2); + + // parallel computation + while (i<=n_4) { + val1 = _mm256_loadu_pd(res+i); + val2 = _mm256_loadu_pd(vec+i); + val1 = _mm256_mul_pd(val1, sclpar1); + val2 = _mm256_mul_pd(val2, sclpar2); + 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] = res[i]*scl1 + vec[i]*scl2; + res[i+1] = res[i+1]*scl1 + vec[i+1]*scl2; + res[i+2] = res[i+2]*scl1 + vec[i+2]*scl2; + } else if (n_i==2) { + res[i] = res[i]*scl1 + vec[i]*scl2; + res[i+1] = res[i+1]*scl1 + vec[i+1]*scl2; + } else if (n_i==1) { + res[i] = res[i]*scl1 + vec[i]*scl2; + } +} + +// return 1 if vec1==vec2, 0 otherwise +static inline +int mju_compare_avx(const int* vec1, const int* vec2, int n) { + int i = 0; + int n_4 = n - 4; + + // vector part + if (n_4>=0) { + __m128i val1, val2, cmp; + + // parallel computation + while (i<=n_4) { + val1 = _mm_loadu_si128((const __m128i*)(vec1+i)); + val2 = _mm_loadu_si128((const __m128i*)(vec2+i)); + cmp = _mm_cmpeq_epi32(val1, val2); + if (_mm_movemask_epi8(cmp)!= 0xFFFF) { + return 0; + } + i += 4; + } + } + + // scalar part + return !memcmp(vec1+i, vec2+i, (n-i)*sizeof(int)); +} + +#endif // defined(__AVX__) && defined(mjUSEDOUBLE) + +#endif // mjUSEPLATFORMSIMD + +#endif // MUJOCO_SRC_ENGINE_ENGINE_UTIL_SPARSE_AVX_H_