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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 12:43:30