(training_state: TrainingState,
key: PRNGKey,
ref_traj: jnp.ndarray,
mask: jnp.ndarray
)
| 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) |
no test coverage detected