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

StableBaselines3中model.learn的total_timesteps参数及耗时疑问

关于StableBaselines3中model.learn(total_timesteps)的疑问解答

1. total_timesteps的正确定义

total_timesteps是训练过程中环境执行的总步数上限,但它的实际生效会受算法自身的采样/更新策略限制,并非严格精确到你设置的数值。

以你使用的PPO为例,它默认有参数n_steps=2048,这个参数控制每次策略更新前必须收集的环境步数(即一次完整rollout的步数)。当你设置total_timesteps=5时,由于PPO必须完成一次完整的rollout才能进行策略更新,所以实际会收集满2048步才会停止第一次迭代,这就是你看到总步数为2048的核心原因。

只有当total_timesteps大于n_steps的整数倍时,训练才会按你设置的总步数逐步停止(比如设置total_timesteps=4096,就会完成两次rollout迭代)。

2. 参数的详细官方文档位置

StableBaselines3的官方文档中,learn()方法的total_timesteps参数说明,可在对应算法的API文档中查看:

  • 对于PPO,直接查看PPO类的learn()方法说明,其中明确标注了该参数是"Total number of timesteps to train for",同时会说明它与n_steps等参数的交互逻辑。
  • 所有算法的基类BaseAlgorithm的learn()方法也有该参数的基础定义,所有基于这个基类的算法(PPO、DQN、SAC等)都遵循相同的核心逻辑。

补充说明测试耗时差异的原因

你手动运行5个episode仅耗时0.000999秒,而model.learn(total_timesteps=5)耗时2.6916秒,差异的核心在于:

  • 手动运行只是单纯执行环境的step和reset操作,没有任何算法相关的计算逻辑。
  • 调用model.learn()时,除了环境步数执行,还包含:
    • 环境的包装处理(Monitor、DummyVecEnv)
    • 策略网络的前向/反向传播计算
    • 采样数据的收集、存储与预处理
    • 策略更新的完整流程
      这些算法本身的计算操作才是耗时的主要来源,而非环境步数本身。

附测试代码(供参考)

自定义环境类

class ShowerEnv(Env):
    listTemperatureKnob = (10, 30, 50)
    shower_length = 3

    def __init__(self):
        self.action_space = Discrete(3)
        self.observation_space = Box(low=np.array([0], dtype=np.float32), high=np.array([100], dtype=np.float32))
        self.reset()
        
    def step(self, action):
        temperatureShowerHead = self.listTemperatureKnob[action]
        self.state = temperatureShowerHead
        self.shower_length -= 1 

        if self.state > 27 and self.state < 33: 
            reward =1 
        else: 
            reward = -1 
        
        if self.shower_length <= 0: 
            done = True
        else:
            done = False
        
        # Apply temperature noise
        #self.state += random.randint(-1,1)
        # Set placeholder for info
        info = {"temperatureShower":temperatureShowerHead}

        obs = np.array([self.state], dtype=np.float32)
        # Return step information
        return obs, reward, done, info

    def render(self):
        pass
    
    def reset(self):
        self.state = self.listTemperatureKnob[0]
        self.shower_length = 3
        obs = np.array([self.state], dtype=np.float32)

        return obs

env = ShowerEnv()
check_env(env, warn=True)

手动运行5个episode

timeStart = time.time()

episodes = 5
for episode in range(1, episodes+1):
    state = env.reset()
    done = False
    score = 0 
    
    while not done:
        env.render()
        action = env.action_space.sample()
        n_state, reward, done, info = env.step(action)
        score+=reward
    print('Episode:{} Score:{}'.format(episode, score))
env.close()

timeEnd = time.time()
print("Elapsed Time: " + str(timeEnd- timeStart))

耗时:0.000999秒

执行模型训练

log_path = os.path.join('Training', 'Logs')
model = PPO("MlpPolicy", env, verbose=1, tensorboard_log=log_path)

timeStart = time.time()
model.learn(total_timesteps=5)

timeEnd = time.time()
print("Elapsed Time: " + str(timeEnd- timeStart))

输出

Using cpu device
Wrapping the env with a `Monitor` wrapper
Wrapping the env in a DummyVecEnv.
Logging to Training\Logs\PPO_7
---------------------------------
| rollout/           |          |
|    ep_len_mean     | 3        |
|    ep_rew_mean     | -1.14    |
| time/              |          |
|    fps             | 1496     |
|    iterations      | 1        |
|    time_elapsed    | 1        |
|    total_timesteps | 2048     |
---------------------------------

耗时:2.6916秒


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 20:20:55