(self, pipeline_state: base.State)
| 82 | ) |
| 83 | |
| 84 | def _get_obs(self, pipeline_state: base.State) -> jax.Array: |
| 85 | return jnp.concatenate([pipeline_state.q, pipeline_state.qd], axis=-1) |
| 86 | |
| 87 | def _get_reward(self, state) -> jax.Array: |
| 88 | # x_feet = state.pipeline_state.x.pos[self.track_body_idx[-2:], 0] |