发布于2026-07-20 阅读(0)
扫一扫,手机访问
先说几个核心判断:在PyTorch里跑PPO,尤其是处理离散动作空间,很多细节官方并没有帮你封装好。你得自己搭Categorical采样、算log_prob和entropy,还得用GAE算advantage,向量环境里的张量reshape稍不留神就会翻车。entropy_coef也不是随便设个常数就完事的,得按动作数归一化处理。下面把这些关键点逐一拆开聊。

PyTorch本身并没有torch.nn.PPO或者torchrl.algorithms.PPO这种拿来就用的高层封装(注意,torchrl库虽然提供了PPO,但默认用Normal分布处理连续动作,而且对Categorical的支持并不稳定)。如果你要跑离散动作的PPO,就得手动搭建策略网络、写Categorical(logits=...).sample()、再算log_prob和entropy——别指望nn.Module能自动帮你完成策略梯度所需的概率计算链条。
不少新手容易踩坑的地方是:直接用torch.argmax(logits)选动作。这玩意儿不可导,策略梯度根本传不回来;或者忘了在log_prob里把action转成long类型,导致索引越界报错。
logits的输出维度必须和动作空间大小严格一致。比如env.action_space.n == 4,那head层就要输出nn.Linear(..., 4)。Categorical(logits=logits).log_prob(action)把概率值抓下来。注意action必须是long类型,否则控制台会直接甩你一个Expected dtype long的报错。forward()方法里直接给sample()——训练的时候你得留着logits,后面多步loss计算都要用它。采样这一步应该放在训练循环里来做。自己拿multiprocessing去fork多个gym.Env实例,十有八九会遇到卡死或者共享随机种子的坑。gym.vector已经帮你封装好了批量重置、批量step、自动堆叠obs/rew/done这些操作。用AsyncVectorEnv能隐藏一部分IO延迟,但调试起来比较头疼;初学者建议从SyncVectorEnv开始入手。
这里有个关键点需要注意:向量环境返回的obs形状是(num_envs, *obs_shape),喂给网络之前通常要view(-1, *obs_shape)把batch维度压平;但logits输出之后,又得再view(num_envs, -1)把动作维度对齐回来,否则Categorical会把所有环境的动作混在一起采样,结果就全乱套了。
num_envs=8,意味着每轮收集的是8条轨迹片段,而不是单条轨迹重复跑8次。reset()返回的obs已经是torch.Tensor了(如果你设了dtype=torch.float32的话),不需要再额外用torch.from_numpy转换。seed——正确的做法是用env.seed(seed + i)来保证各环境的初始状态不同,否则梯度更新会失效。在离散动作空间下,reward往往很稀疏、方差很大,直接算return = sum(reward[t:])会导致梯度爆炸或者收敛极慢。必须用广义优势估计(GAE)来控制bias-variance的权衡。gamma和gae_lambda是敏感参数,值得认真调:
gamma通常取0.99,短周期任务可以降到0.95。一旦低于0.9,agent就会变得极其短视,只看得见眼前的奖励。gae_lambda取0.95是比较稳妥的起点。设为1.0就退化成Monte Carlo,设为0.0则退化成one-step TD。实际工程中0.90~0.97这个区间最常用。next_nonterminal = 1.0 - done,done是bool数组,必须先转成float再参与乘法运算。compute_advantage()方法里用detach()把value网络的梯度切掉——那是value loss要去处理的事情,advantage本身不应该带梯度。很多开源实现直接把entropy_coef = 0.01写死,但一旦动作数发生变化(比如从4个变成18个),这个固定值会让探索强度剧烈波动。更鲁棒的做法是按动作空间大小归一化:
entropy_coef = 0.01 / np.log(env.single_action_space.n)
如果不这么做,小动作空间(比如只有2个分类)下熵项太弱,策略容易坍缩;大动作空间(比如Atari的18个按键)下熵项又太强,agent会变成乱按一通。
另一个容易忽略的问题是:entropy必须从当前batch的logits重新计算,不能复用旧的log_prob。正确的写法是Categorical(logits=logits).entropy().mean(),而不是用-log_prob.mean()——后者只是负对数似然,并不是真实的熵。
如果你在训练日志里看到entropy项持续下降,但reward就是死活不涨,大概率是entropy_coef设得太小,或者用了错误的entropy计算方式。这个时候不妨从这两个方向入手排查。
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
售后无忧
立即购买>office旗舰店
正版软件
正版软件
正版软件
正版软件
正版软件
1
2
3
7
8