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

Function training_epoch

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

Source from the content-addressed store, hash-verified

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

Callers 1

Calls 2

clip_by_global_normFunction · 0.85
TrainingStateClass · 0.85

Tested by

no test coverage detected