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: