(cvae_params, normalizer_params, key, ref_traj, mask)
| 175 | return (nstate, key), (nstate.reward, env_state.obs, obs_latent, nstate.metrics) |
| 176 | |
| 177 | def loss(cvae_params, normalizer_params, key, ref_traj, mask): |
| 178 | encoder_params, policy_params = cvae_params[0], cvae_params[1] |
| 179 | normalizer_encoder, normalizer_policy = normalizer_params |
| 180 | key_reset, key_scan = jax.random.split(key) |
| 181 | env_state = env.reset_ref(jax.random.split(key_reset, num_envs // process_count), ref_traj, mask) |
| 182 | f = functools.partial( |
| 183 | env_step, |
| 184 | policy=make_policy((normalizer_policy, policy_params)), |
| 185 | encoder=make_encoder((normalizer_encoder, encoder_params), deterministic=deterministic) # todo: normalize |
| 186 | ) |
| 187 | (rewards, |
| 188 | obs, obs_latent, metrics) = jax.lax.scan(f, (env_state, key_scan), |
| 189 | (jnp.array(range(episode_length // action_repeat))))[1] |
| 190 | return -jnp.mean(rewards), (rewards, obs, obs_latent, metrics) |
| 191 | |
| 192 | loss_grad = jax.grad(loss, has_aux=True) |
| 193 |
nothing calls this directly
no test coverage detected