| 9 | |
| 10 | |
| 11 | class Walker2d(PipelineEnv): |
| 12 | |
| 13 | def __init__(self): |
| 14 | path = epath.resource_path("brax") / "envs/assets/walker2d.xml" |
| 15 | sys = mjcf.load(path) |
| 16 | |
| 17 | self._reset_noise_scale = 5e-3 |
| 18 | |
| 19 | super().__init__(sys=sys, backend="positional", n_frames=20) |
| 20 | |
| 21 | def reset(self, rng: jax.Array) -> State: |
| 22 | """Resets the environment to an initial state.""" |
| 23 | rng, rng1, rng2 = jax.random.split(rng, 3) |
| 24 | |
| 25 | low, hi = -self._reset_noise_scale, self._reset_noise_scale |
| 26 | qpos = self.sys.init_q + jax.random.uniform( |
| 27 | rng1, (self.sys.q_size(),), minval=low, maxval=hi |
| 28 | ) |
| 29 | qvel = jax.random.uniform(rng2, (self.sys.qd_size(),), minval=low, maxval=hi) |
| 30 | |
| 31 | pipeline_state = self.pipeline_init(qpos, qvel) |
| 32 | |
| 33 | obs = self._get_obs(pipeline_state) |
| 34 | reward, done = jp.zeros(2) |
| 35 | return State(pipeline_state, obs, reward, done, {}) |
| 36 | |
| 37 | def step(self, state: State, action: jax.Array) -> State: |
| 38 | """Runs one timestep of the environment's dynamics.""" |
| 39 | pipeline_state0 = state.pipeline_state |
| 40 | assert pipeline_state0 is not None |
| 41 | pipeline_state = self.pipeline_step(pipeline_state0, action) |
| 42 | |
| 43 | obs = self._get_obs(pipeline_state) |
| 44 | reward = self._get_reward(pipeline_state) |
| 45 | |
| 46 | return state.replace( |
| 47 | pipeline_state=pipeline_state, obs=obs, reward=reward, done=0.0 |
| 48 | ) |
| 49 | |
| 50 | def _get_obs(self, pipeline_state: base.State) -> jax.Array: |
| 51 | """Returns the environment observations.""" |
| 52 | position = pipeline_state.q |
| 53 | position = position.at[1].set(pipeline_state.x.pos[0, 2]) |
| 54 | velocity = jp.clip(pipeline_state.qd, -10, 10) |
| 55 | |
| 56 | return jp.concatenate((position, velocity)) |
| 57 | |
| 58 | def _get_reward(self, pipeline_state: base.State) -> jax.Array: |
| 59 | """Returns the environment reward.""" |
| 60 | return ( |
| 61 | pipeline_state.x.pos[0, 0] |
| 62 | - jp.clip(jp.abs(pipeline_state.x.pos[0, 2] - 1.1), -1.0, 1.0) * 0.5 |
| 63 | ) |