MCPcopy Create free account
hub / github.com/LeCAR-Lab/model-based-diffusion / PushT

Class PushT

mbd/envs/pushT.py:15–74  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

13
14
15class 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

Callers 2

mainFunction · 0.85
get_envFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected