(
self,
obs_dim_dict,
module_config_dict,
num_actions,
init_noise_std,
)
| 82 | |
| 83 | class ActorCritic(nn.Module): |
| 84 | def __init__( |
| 85 | self, |
| 86 | obs_dim_dict, |
| 87 | module_config_dict, |
| 88 | num_actions, |
| 89 | init_noise_std, |
| 90 | ): |
| 91 | super(ActorCritic, self).__init__() |
| 92 | |
| 93 | self.actor_module = Actor(obs_dim_dict, module_config_dict.actor, num_actions) |
| 94 | |
| 95 | critic_module_config_dict = module_config_dict.critic |
| 96 | self.critic_net_type = critic_module_config_dict.get("type", "MLP") |
| 97 | if self.critic_net_type == "MLP": |
| 98 | self.critic_module = BaseModule(obs_dim_dict, critic_module_config_dict) |
| 99 | else: |
| 100 | raise NotImplementedError |
| 101 | |
| 102 | # Action noise |
| 103 | self.std = nn.Parameter(init_noise_std * torch.ones(num_actions)) |
| 104 | self.fix_sigma = module_config_dict.actor.get("fix_sigma", False) |
| 105 | self.max_sigma = module_config_dict.actor.get("max_sigma", 1.0) |
| 106 | self.min_sigma = module_config_dict.actor.get("min_sigma", 0.1) |
| 107 | |
| 108 | if self.fix_sigma: |
| 109 | self.std.requires_grad = False |
| 110 | self.distribution = None |
| 111 | # disable args validation for speedup |
| 112 | Normal.set_default_validate_args = False |
| 113 | |
| 114 | @property |
| 115 | def actor(self): |
nothing calls this directly
no test coverage detected