Resets the environment to an initial state.
(self, rng: jax.Array)
| 17 | super().__init__(sys=sys, backend="positional", n_frames=7) |
| 18 | |
| 19 | def reset(self, rng: jax.Array) -> State: |
| 20 | """Resets the environment to an initial state.""" |
| 21 | rng, rng1, rng2 = jax.random.split(rng, 3) |
| 22 | |
| 23 | low, hi = -0.01, 0.01 |
| 24 | qpos = self.sys.init_q + jax.random.uniform( |
| 25 | rng1, (self.sys.q_size(),), minval=-0.01, maxval=0.01 |
| 26 | ) |
| 27 | qvel = jax.random.uniform(rng2, (self.sys.qd_size(),), minval=low, maxval=hi) |
| 28 | |
| 29 | pipeline_state = self.pipeline_init(qpos, qvel) |
| 30 | obs = self._get_obs(pipeline_state, jnp.zeros(self.sys.act_size())) |
| 31 | reward, done = jnp.zeros(2) |
| 32 | return State(pipeline_state, obs, reward, done, {}) |
| 33 | |
| 34 | def step(self, state: State, action: jax.Array) -> State: |
| 35 | """Runs one timestep of the environment's dynamics.""" |