From 210cf86486440f9c9eb876f0ad431d4073266b74 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Tue, 10 Feb 2026 12:09:37 -0800 Subject: [PATCH] Use kernel cache dir per mjx-warp test fixture. PiperOrigin-RevId: 868259579 Change-Id: I9f1346899591ca603ad2e34949d22b78c464c2ad --- mjx/mujoco/mjx/warp/collision_driver_test.py | 21 ++++++---- mjx/mujoco/mjx/warp/forward_test.py | 42 ++++++++++++-------- mjx/mujoco/mjx/warp/smooth_test.py | 21 ++++++---- 3 files changed, 52 insertions(+), 32 deletions(-) diff --git a/mjx/mujoco/mjx/warp/collision_driver_test.py b/mjx/mujoco/mjx/warp/collision_driver_test.py index 93e72db3..36109eef 100644 --- a/mjx/mujoco/mjx/warp/collision_driver_test.py +++ b/mjx/mujoco/mjx/warp/collision_driver_test.py @@ -39,18 +39,23 @@ _FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1' class CollisionTest(absltest.TestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + if mjxw.WARP_INSTALLED: + cls.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = cls.tempdir.name + + @classmethod + def tearDownClass(cls): + super().tearDownClass() + if hasattr(cls, 'tempdir'): + cls.tempdir.cleanup() + def setUp(self): super().setUp() - if mjxw.WARP_INSTALLED: - 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 fe65576b..aa0700a6 100644 --- a/mjx/mujoco/mjx/warp/forward_test.py +++ b/mjx/mujoco/mjx/warp/forward_test.py @@ -41,18 +41,23 @@ _FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1' class ForwardTest(parameterized.TestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + if mjxw.WARP_INSTALLED: + cls.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = cls.tempdir.name + + @classmethod + def tearDownClass(cls): + super().tearDownClass() + if hasattr(cls, 'tempdir'): + cls.tempdir.cleanup() + def setUp(self): super().setUp() - if mjxw.WARP_INSTALLED: - 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', @@ -198,18 +203,23 @@ class ForwardTest(parameterized.TestCase): class StepTest(parameterized.TestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + if mjxw.WARP_INSTALLED: + cls.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = cls.tempdir.name + + @classmethod + def tearDownClass(cls): + super().tearDownClass() + if hasattr(cls, 'tempdir'): + cls.tempdir.cleanup() + def setUp(self): super().setUp() - if mjxw.WARP_INSTALLED: - 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 b2f867a5..92c02875 100644 --- a/mjx/mujoco/mjx/warp/smooth_test.py +++ b/mjx/mujoco/mjx/warp/smooth_test.py @@ -41,18 +41,23 @@ _FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1' class SmoothTest(parameterized.TestCase): + @classmethod + def setUpClass(cls): + super().setUpClass() + if mjxw.WARP_INSTALLED: + cls.tempdir = tempfile.TemporaryDirectory() + wp.config.kernel_cache_dir = cls.tempdir.name + + @classmethod + def tearDownClass(cls): + super().tearDownClass() + if hasattr(cls, 'tempdir'): + cls.tempdir.cleanup() + def setUp(self): super().setUp() - if mjxw.WARP_INSTALLED: - 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: