def actor_loss_fn(self, logprobs, old_logprobs, advantages, mask): ## policy gradient loss log_ratio = (logprobs - old_logprobs) * mask ratio = torch.exp(log_ratio) pg_loss1 = -advantages * ratio pg_loss2 = -advantages * torch.clamp(ratio, 1.0 - self.cliprange, 1.0 + self.cliprange) pg_loss = torch.sum(torch.max(pg_loss1, pg_loss2) * mask) / mask.sum() return pg_loss def get_advantages_and_returns(self, values, rewards, start): lastgaelam = 0 advantages_reversed = [] length = rewards.size()[-1] for t in reversed(range(start, length)): nextvalues = values[:, t + 1] if t < length - 1 else 0.0 delta = rewards[:, t] + self.gamma * nextvalues - values[:, t] # TD 残差 lastgaelam = delta + self.gamma * self.lam * lastgaelam # GAE advantages_reversed.append(lastgaelam) advantages = torch.stack(advantages_reversed[::-1], dim=1) returns = advantages + values[:, start:] return advantages.detach(), returns def critic_loss_fn(self, values, old_values, returns, mask): ## value loss values_clipped = torch.clamp( values, old_values - self.cliprange_value, old_values + self.cliprange_value, ) vf_loss1 = (values - returns)**2 vf_loss2 = (values_clipped - returns)**2 vf_loss = 0.5 * torch.sum( torch.max(vf_loss1, vf_loss2) * mask) / mask.sum() return vf_loss