| 244 | return action |
| 245 | |
| 246 | def evaluate(self, state): |
| 247 | batch_mu, batch_log_sigma = self.policy_net(state, eval=True) |
| 248 | batch_sigma = torch.exp(batch_log_sigma) |
| 249 | dist = Normal(batch_mu, batch_sigma) |
| 250 | noise = Normal(0, 1) |
| 251 | z = noise.sample() |
| 252 | action = torch.tanh(batch_mu + batch_sigma*z.to(device)) |
| 253 | log_prob = dist.log_prob(batch_mu + batch_sigma * z.to(device)) - torch.log(1 - action.pow(2) + min_Val) |
| 254 | log_prob = log_prob.sum(dim = 2, keepdim=True) / 2 |
| 255 | return action, log_prob, z, batch_mu, batch_log_sigma |
| 256 | |
| 257 | def update(self, replay_buffer, batch_size, current_map_level): |
| 258 | for _ in range(self.args.gradient_steps): |