Small fix to testspeed

PiperOrigin-RevId: 840607453
Change-Id: If913866cc0a3d4d3e6c39c7b5948fa167678ccad
This commit is contained in:
Baruch Tabanpour
2025-12-05 01:27:36 -08:00
committed by Copybara-Service
parent 237a5e5c55
commit fc109f0ab5
+3 -3
View File
@@ -93,7 +93,7 @@ def benchmark(
@jax.vmap
def init(key):
d = mjx.make_data(
m, impl=mx.impl, nconmax=_NCONMAX.value, njmax=_NJMAX.value
m, impl=mx.impl, naconmax=_NACONMAX.value, njmax=_NJMAX.value
)
return d
@@ -262,7 +262,7 @@ def benchmark_raw_jax_warp(
d_ = mujoco.MjData(m)
mw = mjwarp.put_model(m)
d = mjwarp.put_data(
m, d_, nworld=nenv, nconmax=_NCONMAX.value, njmax=_NJMAX.value
m, d_, nworld=nenv, naconmax=_NACONMAX.value, njmax=_NJMAX.value
)
jax_unroll_fn = jax_jit(unroll)
@@ -310,7 +310,7 @@ def benchmark_raw_warp(
mw = mjwarp.put_model(m)
dw = mjwarp.make_data(
m, nworld=nenv, nconmax=_NCONMAX.value, njmax=_NJMAX.value
m, nworld=nenv, nconmax=_NACONMAX.value, njmax=_NJMAX.value
)
if function == 'kinematics':