Brax env with eval metrics.
| 5 | |
| 6 | |
| 7 | class EvalWrapper(wrappers.EvalWrapper): |
| 8 | """Brax env with eval metrics.""" |
| 9 | |
| 10 | def reset_ref(self, rng: jp.ndarray, ref_traj: jp.ndarray, mask: jp.ndarray) -> brax_env.State: |
| 11 | |
| 12 | reset_state = self.env.reset_ref(rng, ref_traj, mask) |
| 13 | reset_state.metrics['reward'] = reset_state.reward |
| 14 | eval_metrics = wrappers.EvalMetrics( |
| 15 | episode_metrics=jax.tree_util.tree_map(jp.zeros_like, |
| 16 | reset_state.metrics), |
| 17 | active_episodes=jp.ones_like(reset_state.reward), |
| 18 | episode_steps=jp.zeros_like(reset_state.reward)) |
| 19 | reset_state.info['eval_metrics'] = eval_metrics |
| 20 | return reset_state |
| 21 | |
| 22 | |
| 23 | class AutoResetWrapper(wrappers.AutoResetWrapper): |
nothing calls this directly
no outgoing calls
no test coverage detected