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

Function training_epoch_with_timing

diffmimic/brax_lib/agent_diffmimic.py:239–258  ·  view source on GitHub ↗
(training_state: TrainingState,
                                   key: PRNGKey,
                                   ref_traj: jnp.ndarray,
                                   mask: jnp.ndarray
                                   )

Source from the content-addressed store, hash-verified

237
238 # Note that this is NOT a pure jittable method.
239 def training_epoch_with_timing(training_state: TrainingState,
240 key: PRNGKey,
241 ref_traj: jnp.ndarray,
242 mask: jnp.ndarray
243 ) -> Tuple[TrainingState, Metrics]:
244 nonlocal training_walltime
245 t = time.time()
246 (training_state, metrics) = training_epoch(training_state, key, ref_traj, mask)
247 metrics = jax.tree_util.tree_map(jnp.mean, metrics)
248 jax.tree_util.tree_map(lambda x: x.block_until_ready(), metrics)
249
250 epoch_training_time = time.time() - t
251 training_walltime += epoch_training_time
252 sps = (episode_length * num_envs) / epoch_training_time * local_devices_to_use
253 metrics = {
254 'training/sps': sps,
255 'training/walltime': training_walltime,
256 **{f'training/{name}': value for name, value in metrics.items()}
257 }
258 return training_state, metrics
259
260 key = jax.random.PRNGKey(seed)
261 global_key, local_key = jax.random.split(key)

Callers 1

trainFunction · 0.85

Calls 1

training_epochFunction · 0.85

Tested by

no test coverage detected