MCPcopy Create free account
hub / github.com/MotrixLab/insactor / Evaluator

Class Evaluator

diffmimic/brax_lib/acting.py:79–153  ·  view source on GitHub ↗

Class to run evaluations.

Source from the content-addressed store, hash-verified

77
78# TODO: Consider moving this to its own file.
79class Evaluator:
80 """Class to run evaluations."""
81
82 def __init__(self, eval_env: envs.Env,
83 eval_policy_fn: Callable[[PolicyParams],
84 Policy],
85 eval_encoder_fn: Callable[[PolicyParams],
86 Policy],
87 num_eval_envs: int,
88 episode_length: int, action_repeat: int, key: PRNGKey):
89 """Init.
90
91 Args:
92 eval_env: Batched environment to run evals on.
93 eval_policy_fn: Function returning the policy from the policy parameters.
94 num_eval_envs: Each env will run 1 episode in parallel for each eval.
95 episode_length: Maximum length of an episode.
96 action_repeat: Number of physics steps per env step.
97 key: RNG key.
98 """
99 self._key = key
100 self._eval_walltime = 0.
101
102 # eval_env = envs.wrappers.EvalWrapper(eval_env)
103 eval_env = wrappers.EvalWrapper(eval_env)
104
105 def generate_eval_unroll(cvae_params: PolicyParams,
106 key: PRNGKey,
107 ref_traj: jnp.ndarray,
108 mask: jnp.ndarray) -> (envs.State, brax.QP):
109 reset_keys = jax.random.split(key, num_eval_envs)
110 # eval_first_state = eval_env.reset(reset_keys)
111 eval_first_state = eval_env.reset_ref(reset_keys, ref_traj, mask)
112 (normalizer_encoder, normalizer_policy), (encoder_params, policy_params) = cvae_params
113 return generate_unroll(
114 eval_env,
115 eval_first_state,
116 eval_policy_fn((normalizer_policy, policy_params)),
117 eval_encoder_fn((normalizer_encoder, encoder_params)),
118 key,
119 unroll_length=episode_length // action_repeat)
120
121 self._generate_eval_unroll = jax.jit(generate_eval_unroll)
122 self._steps_per_unroll = episode_length * num_eval_envs
123
124 def run_evaluation(self,
125 cvae_params: PolicyParams,
126 training_metrics: Metrics,
127 ref_traj: jnp.ndarray,
128 mask: jnp.ndarray,
129 aggregate_episodes: bool = True) -> Metrics:
130 """Run one epoch of evaluation."""
131 self._key, unroll_key = jax.random.split(self._key)
132
133 t = time.time()
134 eval_state, (qp_list, latent_list) = self._generate_eval_unroll(cvae_params, unroll_key, ref_traj, mask)
135 eval_metrics = eval_state.info['eval_metrics']
136 eval_metrics.active_episodes.block_until_ready()

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected