Make encoder networks.
(
observation_size: int,
latent_size: int,
conditional: bool = True,
random: bool = False,
preprocess_observations_fn: types.PreprocessObservationFn = types
.identity_observation_preprocessor,
hidden_layer_sizes: Sequence[int] = (32,) * 4,
activation: networks.ActivationFn = linen.swish)
| 93 | |
| 94 | |
| 95 | def make_encoder_networks( |
| 96 | observation_size: int, |
| 97 | latent_size: int, |
| 98 | conditional: bool = True, |
| 99 | random: bool = False, |
| 100 | preprocess_observations_fn: types.PreprocessObservationFn = types |
| 101 | .identity_observation_preprocessor, |
| 102 | hidden_layer_sizes: Sequence[int] = (32,) * 4, |
| 103 | activation: networks.ActivationFn = linen.swish) -> EncoderNetworks: |
| 104 | """Make encoder networks.""" |
| 105 | parametric_latent_distribution = distribution.NormalTanhDistribution( |
| 106 | event_size=latent_size) |
| 107 | policy_network = networks.make_policy_network( |
| 108 | parametric_latent_distribution.param_size, |
| 109 | observation_size, |
| 110 | preprocess_observations_fn=preprocess_observations_fn, |
| 111 | hidden_layer_sizes=hidden_layer_sizes, activation=activation) |
| 112 | prior_network = networks.make_policy_network( |
| 113 | parametric_latent_distribution.param_size, |
| 114 | observation_size, |
| 115 | preprocess_observations_fn=preprocess_observations_fn, |
| 116 | hidden_layer_sizes=hidden_layer_sizes, activation=activation) |
| 117 | return EncoderNetworks( |
| 118 | policy_network=policy_network, |
| 119 | prior_network=prior_network, |
| 120 | parametric_latent_distribution=parametric_latent_distribution, |
| 121 | conditional=conditional, |
| 122 | random=random) |
nothing calls this directly
no test coverage detected