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

Function loss

diffmimic/brax_lib/agent_diffmimic.py:177–190  ·  view source on GitHub ↗
(cvae_params, normalizer_params, key, ref_traj, mask)

Source from the content-addressed store, hash-verified

175 return (nstate, key), (nstate.reward, env_state.obs, obs_latent, nstate.metrics)
176
177 def loss(cvae_params, normalizer_params, key, ref_traj, mask):
178 encoder_params, policy_params = cvae_params[0], cvae_params[1]
179 normalizer_encoder, normalizer_policy = normalizer_params
180 key_reset, key_scan = jax.random.split(key)
181 env_state = env.reset_ref(jax.random.split(key_reset, num_envs // process_count), ref_traj, mask)
182 f = functools.partial(
183 env_step,
184 policy=make_policy((normalizer_policy, policy_params)),
185 encoder=make_encoder((normalizer_encoder, encoder_params), deterministic=deterministic) # todo: normalize
186 )
187 (rewards,
188 obs, obs_latent, metrics) = jax.lax.scan(f, (env_state, key_scan),
189 (jnp.array(range(episode_length // action_repeat))))[1]
190 return -jnp.mean(rewards), (rewards, obs, obs_latent, metrics)
191
192 loss_grad = jax.grad(loss, has_aux=True)
193

Callers

nothing calls this directly

Calls 2

make_policyFunction · 0.85
reset_refMethod · 0.45

Tested by

no test coverage detected