MCPcopy Create free account
hub / github.com/google-deepmind/dm_control / HooksTracker

Class HooksTracker

dm_control/composer/hooks_test_utils.py:38–234  ·  view source on GitHub ↗

Helper class for tracking call order of callbacks.

Source from the content-addressed store, hash-verified

36
37
38class 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',

Callers 1

setUpMethod · 0.85

Calls

no outgoing calls

Tested by 1

setUpMethod · 0.68

Used in the wild real call sites across dependent graphs

searching dependent graphs…