| 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 | |