Maintains episode step count and sets done at episode end.
| 31 | |
| 32 | |
| 33 | class EpisodeWrapper(wrappers.EpisodeWrapper): |
| 34 | """Maintains episode step count and sets done at episode end.""" |
| 35 | |
| 36 | def reset_ref(self, rng: jp.ndarray, ref_traj: jp.ndarray, mask: jp.ndarray, text_embedding: jp.ndarray) -> brax_env.State: |
| 37 | state = self.env.reset_ref(rng, ref_traj, mask, text_embedding) |
| 38 | state.info['steps'] = jp.zeros(()) |
| 39 | state.info['truncation'] = jp.zeros(()) |
| 40 | return state |
| 41 | |
| 42 | |
| 43 | class VmapWrapper(wrappers.VmapWrapper): |
nothing calls this directly
no outgoing calls
no test coverage detected