(self, pipeline_state: base.State, action: jax.Array)
| 41 | return state.replace(pipeline_state=pipeline_state, obs=obs, reward=reward) |
| 42 | |
| 43 | def _get_obs(self, pipeline_state: base.State, action: jax.Array) -> jax.Array: |
| 44 | return jnp.concatenate([pipeline_state.q, pipeline_state.qd], axis=-1) |
| 45 | |
| 46 | def _get_reward(self, pipeline_state: base.State) -> jax.Array: |
| 47 | return ( |