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
相关产品推荐
相关产品推荐

