def update_policy(self, log_probs, rewards): discounted_rewards = [] for t in range(len(rewards)): Gt = sum(self.gamma ** i * rewards[t + i] for i in range(len(rewards) - t)) discounted_rewards.append(Gt) discounted_rewards = torch.FloatTensor(discounted_rewards).to(device) discounted_rewards = (discounted_rewards - discounted_rewards.mean()) / (discounted_rewards.std() + 1e-9) policy_loss = [] for log_prob, Gt in zip(log_probs, discounted_rewards): policy_loss.append(-log_prob * Gt) self.optimizer.zero_grad() policy_loss = torch.stack(policy_loss).sum() policy_loss.backward() self.optimizer.step()