Add private function mju_isZeroByte to check for byte-wise zero.
PiperOrigin-RevId: 810456282 Change-Id: I877d0f9225e7a783ed7ff12ce54f4ed0a7889793
This commit is contained in:
committed by
Copybara-Service
parent
bb0878417e
commit
d71d42a374
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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));
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user