| 90 | return value |
| 91 | |
| 92 | def evaluate_actions(self, inputs, uav_aoi, uav_snr, uav_compl, uav_tc_compl, action, temporal_hidden_state=None, |
| 93 | mask=None, spatial_hidden_state=None): |
| 94 | if params.use_rnn: |
| 95 | if params.use_spatial_att: |
| 96 | obs_features, _, _ = self.base(inputs, temporal_hidden_state, mask, spatial_hidden_state) |
| 97 | else: |
| 98 | obs_features, _ = self.base(inputs, temporal_hidden_state, mask) |
| 99 | |
| 100 | else: |
| 101 | obs_features = self.base(inputs) |
| 102 | |
| 103 | full_obs_features = torch.cat([obs_features, uav_aoi, uav_snr, uav_compl, uav_tc_compl], dim=1) |
| 104 | value = self.critic(full_obs_features) |
| 105 | dist_dia = self.dist_dia(full_obs_features) |
| 106 | action_log_probs_dia = dist_dia.log_probs(action) |
| 107 | dist_entropy_dia = dist_dia.entropy().mean() |
| 108 | |
| 109 | return value, action_log_probs_dia, dist_entropy_dia |
| 110 | |
| 111 | def print_grad(self): |
| 112 | for name, p in self.base.named_parameters(): |