商城首页欢迎来到中国正版软件门户

您的位置: 首页 > 文章列表 > 编程开发 > Python怎么在PyTorch里跑强化学习PPO算法_Categorical动作采样与环境并行交互

Python怎么在PyTorch里跑强化学习PPO算法_Categorical动作采样与环境并行交互

  发布于2026-07-20 阅读(0)

扫一扫,手机访问

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

Python怎么在PyTorch里跑强化学习PPO算法_Categorical动作采样与环境并行交互

PyTorch里用PPO必须自己实现Categorical采样,官方不提供现成PPO模块

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计算都要用它。采样这一步应该放在训练循环里来做。

环境并行交互要用gym.vector.AsyncVectorEnv或SyncVectorEnv,别手写多进程

自己拿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)来保证各环境的初始状态不同,否则梯度更新会失效。

PPO训练循环里advantage必须用GAE,不能用原始reward-to-go

在离散动作空间下,reward往往很稀疏、方差很大,直接算return = sum(reward[t:])会导致梯度爆炸或者收敛极慢。必须用广义优势估计(GAE)来控制bias-variance的权衡。gammagae_lambda是敏感参数,值得认真调:

  • gamma通常取0.99,短周期任务可以降到0.95。一旦低于0.9,agent就会变得极其短视,只看得见眼前的奖励。
  • gae_lambda0.95是比较稳妥的起点。设为1.0就退化成Monte Carlo,设为0.0则退化成one-step TD。实际工程中0.90~0.97这个区间最常用。
  • GAE公式里有个细节:next_nonterminal = 1.0 - donedonebool数组,必须先转成float再参与乘法运算。
  • 别在compute_advantage()方法里用detach()value网络的梯度切掉——那是value loss要去处理的事情,advantage本身不应该带梯度。

Categorical策略的entropy_loss权重不能硬设为常数

很多开源实现直接把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计算方式。这个时候不妨从这两个方向入手排查。

本文转载于:https://www.php.cn/faq/2323292.html 如有侵犯,请联系zhengruancom@outlook.com删除。
免责声明:正软商城发布此文仅为传递信息,不代表正软商城认同其观点或证实其描述。

热门关注