You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为何在Stable Baselines 3中用CUDA GPU训练CartPole比CPU慢?

CartPole-v1用PPO训练时GPU比CPU慢的原因及解决办法

你遇到的情况不是操作失误,核心问题在于CartPole-v1属于超轻量化的强化学习任务,GPU加速带来的收益完全被额外开销抵消了,具体原因和优化方向如下:

核心原因

  • 模型计算量太小:CartPole默认的MlpPolicy结构非常简单(两层64神经元的全连接层),GPU的并行计算优势需要大规模计算才能体现,这种小模型的计算量根本喂不满GPU,反而要承担CPU与GPU之间的数据传输开销。
  • 环境交互是瓶颈:CartPole的环境步长计算几乎在瞬间完成,且全程在CPU上运行。GPU训练时,每一步都要把观测数据传到GPU、计算后再把结果传回CPU,这个来回传输的时间成本远超过GPU节省的计算时间。
  • 小批量放大开销:Stable Baselines 3中PPO的默认batch size较小,小批量下GPU的调度成本占比极高,进一步拉低了整体效率。

优化建议

  • 提升训练规模:把total_timesteps调到10万甚至更高,当模型训练的计算量占比提升后,GPU的并行优势会逐渐显现。
  • 增大batch size:初始化PPO时设置更大的batch_size(比如batch_size=2048),让GPU的并行计算能力得到充分利用。
  • 使用并行环境:创建多个并行环境来增加数据吞吐量,减少CPU-GPU数据传输的相对开销,示例代码:
from stable_baselines3.common.env_util import make_vec_env
import gym
import time
from stable_baselines3 import PPO

env = make_vec_env("CartPole-v1", n_envs=8)

t1 = time.time()
model = PPO("MlpPolicy", env, verbose=0, device="cuda", batch_size=2048)
model.learn(total_timesteps=100_000)
print(f"Time with cuda : {time.time()-t1:.2f}s")

t1 = time.time()
model = PPO("MlpPolicy", env, verbose=0, device="cpu", batch_size=2048)
model.learn(total_timesteps=100_000)
print(f"Time with cpu : {time.time()-t1:.2f}s")

env.close()
  • 任务匹配硬件:如果只是训练CartPole这类简单任务,直接用CPU更高效;GPU更适合复杂环境(如Atari、MuJoCo)或大模型的训练场景。

你的代码及运行结果

训练代码

import gym
import time
from stable_baselines3 import PPO

env = gym.make("CartPole-v1")

t1 = time.time()
model = PPO("MlpPolicy", env, verbose=0, device="cuda")
model.learn(total_timesteps=10_000)
print(f"Time with cuda : {time.time()-t1:.2f}s")

t1 = time.time()
model = PPO("MlpPolicy", env, verbose=0, device="cpu")
model.learn(total_timesteps=10_000)
print(f"Time with cpu : {time.time()-t1:.2f}s")

env.close()

运行输出

Time with cuda : 21.76s
Time with cpu : 13.33s

内容的提问来源于stack exchange,提问作者Joel

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.31 13:30:44