Vanilla Policy Gradient训练异常:损失下降但Agent未学习
问题现象
在为CartPole-v1环境实现基础Vanilla Policy Gradient(VPG)算法时,出现矛盾的训练现象:损失持续下降(说明模型在更新参数),但回合总奖励(存活步数)不断降低,最终稳定在9-10步(杆倒下的最小步数),相当于模型"学习变得更差"。
核心公式
使用的折扣回报公式:
$$ Q_{k,t} = \sum_{i=0}{\gamma^{i-t} r_i} $$
损失公式:
$$ L = -\sum_{k,t}Q_{k,t}\log\pi_{\theta}(a_t | s_t) $$
原始实现代码
import gymnasium as gym import torch from torch import nn import torch.nn.functional as F from torch.nn.init import xavier_uniform_ import numpy as np GAMMA = 0.99 LEARNING_RATE = 0.001 BATCH_SIZE = 4 DEVICE = torch.device('mps') class XavierLinear(nn.Linear): def __init__(self, in_features: int, out_features: int, bias: bool = True, device=None, dtype=None) -> None: super().__init__(in_features, out_features, bias, device, dtype) xavier_uniform_(self.weight) class VPG(nn.Module): def __init__(self, input_size, output_size): super(VPG, self).__init__() self.net = nn.Sequential( XavierLinear(input_size, 128), nn.ReLU(), XavierLinear(128, output_size), ) def forward(self, x): return F.softmax(self.net(x), dim=0) def run_episode(model, env): obs = env.reset()[0] obs = torch.Tensor(env.reset()[0]).to(DEVICE) te = tr = False rewards, outputs, actions = [], [], [] while not (te or tr): probs = model(obs) action = probs.multinomial(1).item() obs, r, te, tr, _ = env.step(action) obs = torch.Tensor(obs).to(DEVICE) if (te or tr): r = 0 rewards.append(r) outputs.append(probs) actions.append(action) return torch.Tensor(rewards).to(DEVICE), torch.concatenate(outputs).reshape(len(rewards), 2), actions def discount_rewards(rewards): discounted_r = torch.zeros_like(rewards) additive_r = 0 for idx in range(len(rewards)-1, -1, -1): to_add = GAMMA * additive_r additive_r = to_add + rewards[idx] discounted_r[idx] = additive_r return discounted_r.to(DEVICE) def loss_function(discounted_r, probs, actions): logprobs = torch.log(probs) selected = logprobs[range(probs.shape[0]), actions] # discounted_r = (discounted_r - discounted_r.mean()) / discounted_r.std() weighted = selected * discounted_r return -weighted.sum() # The actual training loop: episode_total_reward = 0 batch_losses = torch.Tensor().to(DEVICE) batch_actions = [] batch_disc_r = torch.Tensor().to(DEVICE) batch_probs = torch.Tensor().to(DEVICE) best_ep_reward = 0 losses, ep_total_lenghts = [], [0] episodes = 0 TARGET_REWARD = 100 env = gym.make("CartPole-v1") model = VPG(env.observation_space.shape[0], 2).to(DEVICE) optim = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE) while np.array(ep_total_lenghts)[-100:].mean() < TARGET_REWARD: rewards, probs, actions = run_episode(model, env) discounted_r = discount_rewards(rewards) episode_total_reward = rewards.shape[0] ep_total_lenghts.append(episode_total_reward) episodes += 1 batch_actions += actions batch_disc_r = torch.concatenate([batch_disc_r, discounted_r]) batch_probs = torch.concatenate([batch_probs, probs]) if episodes % BATCH_SIZE == 0: loss = loss_function(batch_disc_r, batch_probs, batch_actions) losses.append(loss.item()) model.zero_grad() loss.backward() optim.step() batch_actions = [] batch_disc_r = torch.Tensor().to(DEVICE) batch_probs = torch.Tensor().to(DEVICE) print(f"Episode {episodes}. Loss: {loss}. Reward: {episode_total_reward}") print(f"Success in {episodes} episodes. Loss: {loss}. Reward: {episode_total_reward}")
前期尝试无效操作
曾尝试调整损失函数符号、修改奖励机制(非终止步为0,终止步为-1)、手动更新权重等,但均无法改变"损失下降但性能恶化"的结果。
有效解决修改
通过以下三处修改解决了问题:
- 调整奖励信号:在
run_episode函数中,将奖励设置为终止时-1,非终止时0:r = -1 if te else 0 - 归一化折扣回报:取消
loss_function中归一化代码的注释,将折扣回报转换为标准化的优势信号:discounted_r = (discounted_r - discounted_r.mean()) / discounted_r.std() - 损失计算改用均值:在
loss_function中返回加权值均值的负数,而非总和的负数:return - weighted.mean()
问题根源分析
原始奖励机制的误导:
原始代码中,非终止步保留CartPole默认的+1奖励,终止步设为0,导致所有折扣回报$Q_{k,t}$均为正数。此时模型的损失优化目标是最大化$\sum Q_{k,t}\log\pi$,但短回合的早期步骤$Q$值远小于长回合的早期步骤$Q$值——这会让模型误以为"早终止的动作更优"(因为小的$Q$乘以$\log\pi$后,总和更小,对应损失更小),最终导致模型主动学习让杆更快倒下。缺乏优势信号归一化:
未归一化的$Q$值尺度不稳定,不同回合的$Q$值差异会导致梯度更新幅度波动剧烈,可能让模型陷入局部最优或错误的更新方向。归一化后将$Q$转换为相对优势(相对于当前batch的平均水平),能稳定梯度,让模型专注于学习"比平均更好的动作"。损失总和的尺度问题:
使用损失总和会导致梯度幅度随batch大小变化,batch越大梯度越大,容易引发模型震荡或过度更新。改用均值后,损失的尺度与batch大小无关,训练过程更稳定。
内容的提问来源于stack exchange,提问作者wildBass

