(training_state: TrainingState, key: PRNGKey, ref_traj: jnp.ndarray, mask: jnp.ndarray)
| 199 | updates) |
| 200 | |
| 201 | def training_epoch(training_state: TrainingState, key: PRNGKey, ref_traj: jnp.ndarray, mask: jnp.ndarray): |
| 202 | key, key_grad = jax.random.split(key) |
| 203 | grad_raw, (rewards, obs, obs_latent, metrics) = loss_grad(training_state.cvae_params, |
| 204 | training_state.normalizer_params, |
| 205 | key_grad, ref_traj, mask) |
| 206 | assert not grad_raw[0] == grad_raw[1] |
| 207 | |
| 208 | grad = clip_by_global_norm(grad_raw) |
| 209 | grad = jax.lax.pmean(grad, axis_name='i') |
| 210 | grad_raw = jax.lax.pmean(grad_raw, axis_name='i') |
| 211 | params_update, optimizer_state = optimizer.update( |
| 212 | grad, training_state.optimizer_state, params=training_state.cvae_params) |
| 213 | cvae_params = optax.apply_updates(training_state.cvae_params, |
| 214 | params_update) |
| 215 | normalizer_encoder, normalizer_policy = training_state.normalizer_params |
| 216 | normalizer_encoder = running_statistics.update( |
| 217 | normalizer_encoder, obs, pmap_axis_name=_PMAP_AXIS_NAME) |
| 218 | normalizer_policy = running_statistics.update( |
| 219 | normalizer_policy, obs_latent, pmap_axis_name=_PMAP_AXIS_NAME) |
| 220 | |
| 221 | normalizer_params = (normalizer_encoder, normalizer_policy) |
| 222 | metrics = { |
| 223 | 'grad_norm': optax.global_norm(grad_raw), |
| 224 | 'params_norm': optax.global_norm(cvae_params), |
| 225 | 'loss': -1 * rewards, |
| 226 | **metrics |
| 227 | } |
| 228 | return TrainingState( |
| 229 | optimizer_state=optimizer_state, |
| 230 | normalizer_params=normalizer_params, |
| 231 | cvae_params=cvae_params |
| 232 | ), metrics |
| 233 | |
| 234 | training_epoch = jax.pmap(training_epoch, axis_name=_PMAP_AXIS_NAME) |
| 235 |
no test coverage detected