Allow dot products with sparse vectors to specify that one vector uses uncompressed memory.
PiperOrigin-RevId: 561951337 Change-Id: I1310d7c09d9a85ad9b43877155844a9d0ce6edab
This commit is contained in:
committed by
Copybara-Service
parent
70c5fa50f2
commit
29aa5e4a41
@@ -30,9 +30,10 @@
|
||||
//------------------------------ sparse operations using avx ---------------------------------------
|
||||
|
||||
// dot-product, first vector is sparse
|
||||
// flg_unc1: is vec1 memory layout uncompressed
|
||||
static inline
|
||||
mjtNum mju_dotSparse_avx(const mjtNum* vec1, const mjtNum* vec2,
|
||||
const int nnz1, const int* ind1) {
|
||||
const int nnz1, const int* ind1, int flg_unc1) {
|
||||
int i = 0;
|
||||
mjtNum res = 0;
|
||||
int nnz1_4 = nnz1 - 4;
|
||||
@@ -47,20 +48,43 @@ mjtNum mju_dotSparse_avx(const mjtNum* vec1, const mjtNum* vec2,
|
||||
vec2[ind1[2]],
|
||||
vec2[ind1[1]],
|
||||
vec2[ind1[0]]);
|
||||
val1 = _mm256_loadu_pd(vec1);
|
||||
if (flg_unc1) {
|
||||
val1 = _mm256_set_pd(vec1[ind1[3]],
|
||||
vec1[ind1[2]],
|
||||
vec1[ind1[1]],
|
||||
vec1[ind1[0]]);
|
||||
} else {
|
||||
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;
|
||||
if (flg_unc1) {
|
||||
while (i<=nnz1_4) {
|
||||
val1 = _mm256_set_pd(vec1[ind1[i+3]],
|
||||
vec1[ind1[i+2]],
|
||||
vec1[ind1[i+1]],
|
||||
vec1[ind1[i+0]]);
|
||||
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;
|
||||
}
|
||||
} else {
|
||||
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
|
||||
@@ -72,8 +96,14 @@ mjtNum mju_dotSparse_avx(const mjtNum* vec1, const mjtNum* vec2,
|
||||
}
|
||||
|
||||
// scalar part
|
||||
for (; i<nnz1; i++) {
|
||||
res += vec1[i] * vec2[ind1[i]];
|
||||
if (flg_unc1) {
|
||||
for (; i < nnz1; i++) {
|
||||
res += vec1[ind1[i]] * vec2[ind1[i]];
|
||||
}
|
||||
} else {
|
||||
for (; i < nnz1; i++) {
|
||||
res += vec1[i] * vec2[ind1[i]];
|
||||
}
|
||||
}
|
||||
|
||||
return res;
|
||||
@@ -179,7 +209,7 @@ void mju_mulMatVecSparse_avx(mjtNum* res, const mjtNum* mat, const mjtNum* vec,
|
||||
if (!rowsuper) {
|
||||
// regular sparse dot-product
|
||||
for (int r=0; r<nr; r++) {
|
||||
res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]);
|
||||
res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r], /*flg_unc2=*/0);
|
||||
}
|
||||
|
||||
return;
|
||||
@@ -202,7 +232,7 @@ void mju_mulMatVecSparse_avx(mjtNum* res, const mjtNum* mat, const mjtNum* vec,
|
||||
|
||||
// handle remaining rows
|
||||
while (rs>0) {
|
||||
res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]);
|
||||
res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r], /*flg_unc2=*/0);
|
||||
|
||||
r++;
|
||||
rs--;
|
||||
@@ -213,7 +243,7 @@ void mju_mulMatVecSparse_avx(mjtNum* res, const mjtNum* mat, const mjtNum* vec,
|
||||
}
|
||||
|
||||
else {
|
||||
res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]);
|
||||
res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r], /*flg_unc2=*/0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user