Helper class for tracking call order of callbacks.
| 36 | |
| 37 | |
| 38 | class HooksTracker: |
| 39 | """Helper class for tracking call order of callbacks.""" |
| 40 | |
| 41 | def __init__(self, test_case, physics_timestep, control_timestep, |
| 42 | *args, **kwargs): |
| 43 | super().__init__(*args, **kwargs) |
| 44 | self.tracked = False |
| 45 | self._test_case = test_case |
| 46 | self._call_count = collections.defaultdict(lambda: 0) |
| 47 | self._physics_timestep = physics_timestep |
| 48 | self._physics_steps_per_control_step = ( |
| 49 | round(int(control_timestep / physics_timestep))) |
| 50 | |
| 51 | mro = inspect.getmro(type(self)) |
| 52 | self._has_super = mro[mro.index(HooksTracker) + 1] != object |
| 53 | |
| 54 | def assertEqual(self, actual, expected, msg=''): |
| 55 | msg = '{}: {}: {!r} != {!r}'.format(type(self), msg, actual, expected) |
| 56 | self._test_case.assertEqual(actual, expected, msg) |
| 57 | |
| 58 | def assertHooksNotCalled(self, *hook_names): |
| 59 | for hook_name in hook_names: |
| 60 | self.assertEqual( |
| 61 | self._call_count[hook_name], 0, |
| 62 | 'assertHooksNotCalled: hook_name = {!r}'.format(hook_name)) |
| 63 | |
| 64 | def assertHooksCalledOnce(self, *hook_names): |
| 65 | for hook_name in hook_names: |
| 66 | self.assertEqual( |
| 67 | self._call_count[hook_name], 1, |
| 68 | 'assertHooksCalledOnce: hook_name = {!r}'.format(hook_name)) |
| 69 | |
| 70 | def assertCompleteEpisode(self, control_steps): |
| 71 | self.assertHooksCalledOnce('initialize_episode_mjcf', |
| 72 | 'after_compile', |
| 73 | 'initialize_episode') |
| 74 | physics_steps = control_steps * self._physics_steps_per_control_step |
| 75 | self.assertEqual(self._call_count['before_step'], control_steps) |
| 76 | self.assertEqual(self._call_count['before_substep'], physics_steps) |
| 77 | self.assertEqual(self._call_count['after_substep'], physics_steps) |
| 78 | self.assertEqual(self._call_count['after_step'], control_steps) |
| 79 | |
| 80 | def assertPhysicsStepCountEqual(self, physics, expected_count): |
| 81 | actual_count = int(round(physics.time() / self._physics_timestep)) |
| 82 | self.assertEqual(actual_count, expected_count) |
| 83 | |
| 84 | def reset_call_counts(self): |
| 85 | self._call_count = collections.defaultdict(lambda: 0) |
| 86 | |
| 87 | def initialize_episode_mjcf(self, random_state): |
| 88 | """Implements `initialize_episode_mjcf` Composer callback.""" |
| 89 | if self._has_super: |
| 90 | super().initialize_episode_mjcf(random_state) |
| 91 | if not self.tracked: |
| 92 | return |
| 93 | self.assertHooksNotCalled('after_compile', |
| 94 | 'initialize_episode', |
| 95 | 'before_step', |