如何在Stable Baselines 3自定义日志器中获取训练指标(如entropy_loss)
在Stable Baselines 3自定义回调中获取模型训练指标(如entropy_loss)
对于A2C这类on-policy算法,训练损失(包括entropy_loss)是在每次rollout收集完数据后计算的,因此你需要在回调的_on_rollout_end方法中获取这些指标,而非_on_step(_on_step是每一步环境交互后触发,此时还未计算训练损失)。
方法1:通过self.locals直接获取训练变量
Stable Baselines 3执行训练逻辑时,会把训练过程中的关键变量(比如各类损失值)存入self.locals字典,你可以直接从中读取:
import gym from stable_baselines3 import A2C from stable_baselines3.common.callbacks import BaseCallback class MetricLogger(BaseCallback): def __init__(self, verbose=0): super(MetricLogger, self).__init__(verbose) def _on_rollout_end(self) -> None: # 从locals中直接获取entropy_loss entropy_loss = self.locals['entropy_loss'] # 添加你的后续处理逻辑,比如打印、自定义存储等 print(f"当前entropy_loss: {entropy_loss.item()}") env = gym.make('CartPole-v1') model = A2C('MlpPolicy', env, verbose=1) model.learn(total_timesteps=1000, callback=MetricLogger())
方法2:通过self.logger获取已记录的指标
当你设置verbose=1时,控制台打印的指标其实是由SB3内置logger记录的,你可以从self.logger.name_to_value字典中读取所有已记录的指标:
import gym from stable_baselines3 import A2C from stable_baselines3.common.callbacks import BaseCallback class MetricLogger(BaseCallback): def __init__(self, verbose=0): super(MetricLogger, self).__init__(verbose) def _on_rollout_end(self) -> None: # 从logger的已记录指标中获取entropy_loss entropy_loss = self.logger.name_to_value['train/entropy_loss'] print(f"当前entropy_loss: {entropy_loss}") env = gym.make('CartPole-v1') model = A2C('MlpPolicy', env, verbose=1) model.learn(total_timesteps=1000, callback=MetricLogger())
注意事项
- 不同算法的
locals变量名可能略有差异,需对应算法的训练逻辑确认变量名; _on_rollout_end的触发时机与模型训练步骤同步,适合获取训练相关损失指标;_on_step则更适合获取每一步的环境交互数据(如奖励、状态)。
内容的提问来源于stack exchange,提问作者rich
相关产品推荐
相关产品推荐

