| 61 | |
| 62 | |
| 63 | def finish_episode(): |
| 64 | R = 0 |
| 65 | policy_loss = [] |
| 66 | returns = deque() |
| 67 | for r in policy.rewards[::-1]: |
| 68 | R = r + args.gamma * R |
| 69 | returns.appendleft(R) |
| 70 | returns = torch.tensor(returns) |
| 71 | returns = (returns - returns.mean()) / (returns.std() + eps) |
| 72 | for log_prob, R in zip(policy.saved_log_probs, returns): |
| 73 | policy_loss.append(-log_prob * R) |
| 74 | optimizer.zero_grad() |
| 75 | policy_loss = torch.cat(policy_loss).sum() |
| 76 | policy_loss.backward() |
| 77 | optimizer.step() |
| 78 | del policy.rewards[:] |
| 79 | del policy.saved_log_probs[:] |
| 80 | |
| 81 | |
| 82 | def main(): |