(self, rng: jnp.ndarray)
| 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) |
no test coverage detected