(rollout_traj, num_eval_envs, perturb=False)
| 70 | |
| 71 | |
| 72 | def execute_actions(rollout_traj, num_eval_envs, perturb=False): |
| 73 | def make_evaluator(env): |
| 74 | ep_len_eval = 196 |
| 75 | |
| 76 | eval_env = wrappers.EpisodeWrapper(env, ep_len_eval - 1, 1) |
| 77 | eval_env = wrappers.VmapWrapper(eval_env) |
| 78 | eval_env = wrappers.AutoResetWrapper(eval_env) |
| 79 | return acting.Evaluator( |
| 80 | eval_env, |
| 81 | eval_policy_fn=functools.partial(make_policy, deterministic=False), |
| 82 | eval_encoder_fn=functools.partial(make_encoder, deterministic=skip_encoder), |
| 83 | num_eval_envs=num_eval_envs, # Change me!!!!, |
| 84 | episode_length=ep_len_eval, |
| 85 | action_repeat=1, |
| 86 | key=jax.random.PRNGKey(999) |
| 87 | ) |
| 88 | |
| 89 | evaluator_global = make_evaluator(env_global) |
| 90 | evaluator_global_perturb = make_evaluator(env_global_perturb) |
| 91 | |
| 92 | evaluator = evaluator_global_perturb if perturb else evaluator_global |
| 93 | metrics, (qp_list, latent_list) = evaluator.run_evaluation( |
| 94 | params_global, |
| 95 | ref_traj=rollout_traj, |
| 96 | mask=np.ones(rollout_traj.shape[:-1]), |
| 97 | training_metrics={}) |
| 98 | |
| 99 | return serialize_qp(qp_list).transpose(1,0,2)[:, :rollout_traj.shape[1]] |
| 100 | |
| 101 | if __name__ == '__main__': |
| 102 | rollout_traj = np.zeros([16, 120, 247]) |
no test coverage detected