(step_env, state, us)
| 12 | return rews |
| 13 | |
| 14 | def rollout_us(step_env, state, us): |
| 15 | def step(state, u): |
| 16 | state = step_env(state, u) |
| 17 | return state, (state.reward, state.pipeline_state) |
| 18 | |
| 19 | _, (rews, pipline_states) = jax.lax.scan(step, state, us) |
| 20 | return rews, pipline_states |
| 21 | |
| 22 | |
| 23 | def render_us(step_env, sys, state, us): |