| 13 | |
| 14 | |
| 15 | class PushT(PipelineEnv): |
| 16 | def __init__(self, backend: str = "generalized"): |
| 17 | # get path of mbd |
| 18 | sys = mjcf.load(f"{mbd.__path__[0]}/assets/pushT.xml") |
| 19 | |
| 20 | super().__init__(sys, backend=backend, n_frames=5) |
| 21 | |
| 22 | def reset(self, rng: jnp.ndarray) -> State: |
| 23 | rng, rng_goal_xy = jax.random.split(rng) |
| 24 | |
| 25 | q = self.sys.init_q |
| 26 | q = q.at[:2].set(jnp.array([0.1, -0.15])) |
| 27 | q = q.at[5:].set( |
| 28 | jax.random.uniform(rng_goal_xy, (3,), minval=-1.0, maxval=1.0) |
| 29 | * jnp.array([0.2, 0.2, jnp.pi/4]) + jnp.array([-0.4, 0.4, jnp.pi]) |
| 30 | ) |
| 31 | qd = jnp.zeros(self.sys.qd_size()) |
| 32 | pipeline_state = self.pipeline_init(q, qd) |
| 33 | obs = self._get_obs(pipeline_state) |
| 34 | reward = self._get_reward(pipeline_state) |
| 35 | done = self._get_done(pipeline_state) |
| 36 | metrics = {} |
| 37 | return State(pipeline_state, obs, reward, done, metrics) |
| 38 | |
| 39 | def step(self, state: State, action: jnp.ndarray) -> State: |
| 40 | pipeline_state = self.pipeline_step(state.pipeline_state, action) |
| 41 | obs = self._get_obs(pipeline_state) |
| 42 | reward = self._get_reward(pipeline_state) |
| 43 | done = self._get_done(pipeline_state) |
| 44 | return state.replace( |
| 45 | pipeline_state=pipeline_state, obs=obs, reward=reward, done=done |
| 46 | ) |
| 47 | |
| 48 | def _get_obs(self, pipeline_state: pipeline.State) -> jnp.ndarray: |
| 49 | return jnp.concat([pipeline_state.q, pipeline_state.qd], axis=-1) |
| 50 | |
| 51 | def _get_reward(self, pipeline_state: pipeline.State) -> jnp.ndarray: |
| 52 | r_goal = pipeline_state.q[5:7] |
| 53 | r_slider = pipeline_state.q[2:4] |
| 54 | r_pusher = pipeline_state.q[0:2] |
| 55 | theta_goal = pipeline_state.q[7] |
| 56 | theta_slider = pipeline_state.q[4] |
| 57 | d_pusher2slider = jnp.maximum(jnp.linalg.norm(r_pusher - r_slider) - 0.2, 0.0) |
| 58 | return 1.0 - ( |
| 59 | jnp.linalg.norm(r_goal - r_slider) |
| 60 | + (jnp.abs(theta_goal - theta_slider) / jnp.pi) |
| 61 | + d_pusher2slider |
| 62 | ) |
| 63 | |
| 64 | def _get_done(self, pipeline_state: pipeline.State) -> jnp.ndarray: |
| 65 | done = self._get_reward(pipeline_state) > 0.95 |
| 66 | return done.astype(jnp.float32) |
| 67 | |
| 68 | @property |
| 69 | def action_size(self): |
| 70 | return 2 |
| 71 | |
| 72 | @property |