Add private function mju_isZeroByte to check for byte-wise zero.

PiperOrigin-RevId: 810456282
Change-Id: I877d0f9225e7a783ed7ff12ce54f4ed0a7889793
This commit is contained in:
Yuval Tassa
2025-09-23 08:48:43 -07:00
committed by Copybara-Service
parent bb0878417e
commit d71d42a374
8 changed files with 128 additions and 6 deletions
+1 -1
View File
@@ -3389,7 +3389,7 @@ int mju_str2Type(const char* str);
const char* mju_writeNumBytes(size_t nbytes);
const char* mju_warningText(int warning, size_t info);
int mju_isBad(mjtNum x);
int mju_isZero(mjtNum* vec, int n);
int mju_isZero(const mjtNum* vec, int n);
mjtNum mju_standardNormal(mjtNum* num2);
void mju_f2n(mjtNum* res, const float* vec, int n);
void mju_n2f(float* res, const mjtNum* vec, int n);
+1 -1
View File
@@ -1303,7 +1303,7 @@ MJAPI const char* mju_warningText(int warning, size_t info);
MJAPI int mju_isBad(mjtNum x);
// Return 1 if all elements are 0.
MJAPI int mju_isZero(mjtNum* vec, int n);
MJAPI int mju_isZero(const mjtNum* vec, int n);
// Standard normal random number generator (optional second number).
MJAPI mjtNum mju_standardNormal(mjtNum* num2);
+1 -1
View File
@@ -8343,7 +8343,7 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
FunctionParameterDecl(
name='vec',
type=PointerType(
inner_type=ValueType(name='mjtNum'),
inner_type=ValueType(name='mjtNum', is_const=True),
),
),
FunctionParameterDecl(
+27 -1
View File
@@ -1350,7 +1350,7 @@ int mju_isBad(mjtNum x) {
// return 1 if all elements are 0
int mju_isZero(mjtNum* vec, int n) {
int mju_isZero(const mjtNum* vec, int n) {
for (int i=0; i < n; i++) {
if (vec[i] != 0) {
return 0;
@@ -1362,6 +1362,32 @@ int mju_isZero(mjtNum* vec, int n) {
// return 1 if all elements are 0
int mju_isZeroByte(const unsigned char* vec, int n) {
size_t i = 0;
// unroll using 8-byte chunks
const uint64_t* vec64 = (const uint64_t*)vec;
size_t n64 = n / sizeof(uint64_t);
for (; i < n64; ++i) {
if (vec64[i]) {
return 0;
}
}
// remaining bytes
i *= sizeof(uint64_t);
for (; i < n; ++i) {
if (vec[i]) {
return 0;
}
}
return 1;
}
// set integer vector to 0
void mju_zeroInt(int* res, int n) {
memset(res, 0, n*sizeof(int));
+5 -2
View File
@@ -133,8 +133,11 @@ MJAPI const char* mju_warningText(int warning, size_t info);
// return 1 if nan or abs(x)>mjMAXVAL, 0 otherwise
MJAPI int mju_isBad(mjtNum x);
// return 1 if all elements are 0
MJAPI int mju_isZero(mjtNum* vec, int n);
// return 1 if all elements are numerically 0 (-0.0 treated as zero)
MJAPI int mju_isZero(const mjtNum* vec, int n);
// return 1 if all elements are 0x00, can be ~2x faster than mju_isZero
MJAPI int mju_isZeroByte(const unsigned char* vec, int n);
// set integer vector to 0
MJAPI void mju_zeroInt(int* res, int n);
+6
View File
@@ -55,6 +55,12 @@ mujoco_test(
ADDITIONAL_LINK_LIBRARIES benchmark::benchmark absl::core_headers
)
mujoco_test(
iszero_benchmark_test
MAIN_TARGET benchmark::benchmark_main
ADDITIONAL_LINK_LIBRARIES benchmark::benchmark absl::core_headers
)
mujoco_test(
solveLD_benchmark_test
MAIN_TARGET benchmark::benchmark_main
+76
View File
@@ -0,0 +1,76 @@
// Copyright 2025 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.
#include <vector>
#include "benchmark/benchmark.h"
#include <absl/base/attributes.h>
#include <mujoco/mujoco.h>
#include "src/engine/engine_util_misc.h"
namespace mujoco {
namespace {
// Run benchmark for the isZero functions, for mjtNum directly or bytewise.
ABSL_ATTRIBUTE_NO_TAIL_CALL static void IzZeroBenchmark(
benchmark::State& state, int n, bool bytewise) {
std::vector<mjtNum> data(n, 0);
data[n-1] = 1;
if (bytewise) {
for (auto s : state) {
int iszero =
mju_isZeroByte((const unsigned char*)data.data(), n * sizeof(mjtNum));
benchmark::DoNotOptimize(iszero);
}
} else {
for (auto s : state) {
int iszero = mju_isZero(data.data(), n);
benchmark::DoNotOptimize(iszero);
}
}
state.counters["items/s"] =
benchmark::Counter(state.iterations() * n, benchmark::Counter::kIsRate);
}
// Define benchmarks.
void BM_isZeroByte_1e2(benchmark::State& state) {
IzZeroBenchmark(state, 100, true);
}
BENCHMARK(BM_isZeroByte_1e2);
void BM_isZeroByte_1e5(benchmark::State& state) {
IzZeroBenchmark(state, 100000, true);
}
BENCHMARK(BM_isZeroByte_1e5);
void BM_isZero_1e2(benchmark::State& state) {
IzZeroBenchmark(state, 100, false);
}
BENCHMARK(BM_isZero_1e2);
void BM_isZero_1e5(benchmark::State& state) {
IzZeroBenchmark(state, 100000, false);
}
BENCHMARK(BM_isZero_1e5);
} // namespace
} // namespace mujoco
int main(int argc, char** argv) {
benchmark::Initialize(&argc, argv);
benchmark::RunSpecifiedBenchmarks();
return 0;
}
+11
View File
@@ -349,6 +349,17 @@ TEST_F(UtilMiscTest, MjuSparseLower2SymMapPartial) {
EXPECT_THAT(AsVector(mat_res, res_nnz), ElementsAre(1, 2, 0, 2, 3, 0, 6));
}
TEST_F(UtilMiscTest, MjuIsZero) {
mjtNum vec[1] = {1};
EXPECT_EQ(mju_isZero(vec, 1), 0);
EXPECT_EQ(mju_isZero(vec, 0), 1);
vec[0] = 0;
EXPECT_EQ(mju_isZero(vec, 1), 1);
vec[0] = -0.0;
EXPECT_EQ(mju_isZero(vec, 1), 1);
EXPECT_EQ(mju_isZeroByte((const unsigned char*)vec, sizeof(mjtNum)), 0);
}
// --------------------------------- Interpolation -----------------------------
using InterpolationTest = MujocoTest;