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

Method reset

mbd/envs/pushT.py:22–37  ·  view source on GitHub ↗
(self, rng: jnp.ndarray)

Source from the content-addressed store, hash-verified

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)

Callers 1

vis_diffusion.pyFile · 0.45

Calls 4

_get_obsMethod · 0.95
_get_rewardMethod · 0.95
_get_doneMethod · 0.95
StateClass · 0.85

Tested by

no test coverage detected