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

Function make_policy

diffmimic/brax_lib/encoder.py:42–90  ·  view source on GitHub ↗
(params: types.PolicyParams,
                  deterministic: bool = False)

Source from the content-addressed store, hash-verified

40 """Creates params and inference function for the encoder agent."""
41
42 def make_policy(params: types.PolicyParams,
43 deterministic: bool = False) -> types.Policy:
44
45 def policy(observations: types.Observation,
46 key_sample: PRNGKey) -> Tuple[types.Action, types.Extra]:
47 normalizer, (param_p, param_q) = params
48 observations_prev, observations_curr = jnp.split(observations, 2, axis=-1)
49 observation_p_padded = jnp.concatenate([observations_prev, jnp.zeros_like(observations_prev)], axis=-1)
50 logits_p = encoder_networks.prior_network.apply(normalizer, param_p, observation_p_padded)
51 logits_q = encoder_networks.policy_network.apply(normalizer, param_q, observations)
52
53 if encoder_networks.conditional:
54
55 loc_p, _ = jnp.split(logits_p, 2, axis=-1)
56 loc_delta, scale_full = jnp.split(logits_q, 2, axis=-1)
57
58 if not encoder_networks.random:
59 loc = loc_p + loc_delta
60 else:
61 loc = loc_p + loc_delta * 0.
62 scale = jnp.ones_like(scale_full) * scale_full[..., :1]
63
64 logits = jnp.concatenate([loc, scale], axis=-1)
65
66 sigma = jax.nn.softplus(scale) + 0.001
67 kl_div = 0.5 * loc_delta ** 2 / sigma ** 2
68 kl_div = kl_div.mean(-1)
69 if deterministic:
70 return loc, kl_div
71 return encoder_networks.parametric_latent_distribution.sample_no_postprocessing(
72 logits, key_sample), kl_div
73
74 else:
75 loc, scale = jnp.split(logits_q, 2, axis=-1)
76 sigma = jax.nn.softplus(scale) + 0.001
77 log_var = 2 * jnp.log(sigma)
78 kl_div = 0.5 * (jnp.exp(log_var) + loc ** 2 - 1. - log_var)
79 kl_div = kl_div.mean(-1)
80 if deterministic:
81 if observations_curr.shape[-1] == loc.shape[-1]: # skip encoder
82 return observations_curr, kl_div
83 else:
84 return loc, kl_div
85 elif encoder_networks.random:
86 return jax.random.normal(key_sample, shape=loc.shape), kl_div
87 return encoder_networks.parametric_latent_distribution.sample_no_postprocessing(
88 logits_q, key_sample), kl_div
89
90 return policy
91
92 return make_policy
93

Callers 1

lossFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected