From 45476483f60fa327416779a4af6bfa5d06803526 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Thu, 7 Aug 2025 13:27:48 -0700 Subject: [PATCH] Fix race condition in warp tests. PiperOrigin-RevId: 792285453 Change-Id: Ic7616054025ca5c1211ef49e408ee283d8082796 --- mjx/mujoco/mjx/warp/collision_driver_test.py | 9 ++++++++- mjx/mujoco/mjx/warp/forward_test.py | 17 +++++++++++++++-- mjx/mujoco/mjx/warp/smooth_test.py | 9 ++++++++- 3 files changed, 31 insertions(+), 4 deletions(-) diff --git a/mjx/mujoco/mjx/warp/collision_driver_test.py b/mjx/mujoco/mjx/warp/collision_driver_test.py index 4a9263b4..93e72db3 100644 --- a/mjx/mujoco/mjx/warp/collision_driver_test.py +++ b/mjx/mujoco/mjx/warp/collision_driver_test.py @@ -14,6 +14,7 @@ # ============================================================================== """Tests for collision driver.""" import os +import tempfile from absl.testing import absltest import jax @@ -41,9 +42,15 @@ class CollisionTest(absltest.TestCase): def setUp(self): super().setUp() if mjxw.WARP_INSTALLED: - wp.clear_kernel_cache() + self.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = self.tempdir.name np.random.seed(0) + def tearDown(self): + super().tearDown() + if hasattr(self, 'tempdir'): + self.tempdir.cleanup() + _SPHERE_SPHERE = """ diff --git a/mjx/mujoco/mjx/warp/forward_test.py b/mjx/mujoco/mjx/warp/forward_test.py index f1b7b3a3..5fae3d22 100644 --- a/mjx/mujoco/mjx/warp/forward_test.py +++ b/mjx/mujoco/mjx/warp/forward_test.py @@ -16,6 +16,7 @@ import functools import os +import tempfile from absl.testing import absltest from absl.testing import parameterized @@ -43,9 +44,15 @@ class ForwardTest(parameterized.TestCase): def setUp(self): super().setUp() if mjxw.WARP_INSTALLED: - wp.clear_kernel_cache() + self.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = self.tempdir.name np.random.seed(0) + def tearDown(self): + super().tearDown() + if hasattr(self, 'tempdir'): + self.tempdir.cleanup() + @parameterized.parameters( 'pendula.xml', 'humanoid/humanoid.xml', @@ -203,9 +210,15 @@ class StepTest(parameterized.TestCase): def setUp(self): super().setUp() if mjxw.WARP_INSTALLED: - wp.clear_kernel_cache() + self.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = self.tempdir.name np.random.seed(0) + def tearDown(self): + super().tearDown() + if hasattr(self, 'tempdir'): + self.tempdir.cleanup() + @parameterized.product( xml=( 'humanoid/humanoid.xml', diff --git a/mjx/mujoco/mjx/warp/smooth_test.py b/mjx/mujoco/mjx/warp/smooth_test.py index 930bbd47..6d9bc349 100644 --- a/mjx/mujoco/mjx/warp/smooth_test.py +++ b/mjx/mujoco/mjx/warp/smooth_test.py @@ -16,6 +16,7 @@ import functools import os +import tempfile from absl.testing import absltest import jax @@ -42,9 +43,15 @@ class SmoothTest(absltest.TestCase): def setUp(self): super().setUp() if mjxw.WARP_INSTALLED: - wp.clear_kernel_cache() + self.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = self.tempdir.name np.random.seed(0) + def tearDown(self): + super().tearDown() + if hasattr(self, 'tempdir'): + self.tempdir.cleanup() + def test_kinematics(self): """Tests kinematics with unbatched data.""" if not _FORCE_TEST: