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

在Stable-Baselines3自定义回调中访问A2C指标与环境状态

问题解答

一、获取环境的truncated/terminated状态,实现截断终止训练与记录episode时长

完全可以在自定义回调的_on_step方法中获取这两个状态,具体实现如下:

  1. 先在回调类中维护一个变量,用于记录当前episode的步数;
  2. 通过self.locals['infos']获取每个环境的信息字典,其中包含terminated和truncated布尔字段,对应环境的终止/截断状态;
  3. 当检测到环境进入truncated状态时,直接设置self.model.stop_training = True终止训练,同时记录当前episode的时长。

示例代码:

from stable_baselines3.common.callbacks import BaseCallback

class CustomCallback(BaseCallback):
    def __init__(self, verbose=0):
        super().__init__(verbose)
        self.current_episode_steps = 0

    def _on_step(self) -> bool:
        self.current_episode_steps += 1
        infos = self.locals['infos']
        
        for info in infos:
            # 记录episode时长并重置计数
            if info.get('terminated', False) or info.get('truncated', False):
                print(f"Episode ended, duration: {self.current_episode_steps} steps")
                self.current_episode_steps = 0
                
                # 环境截断时终止训练
                if info.get('truncated', False):
                    self.model.stop_training = True
                    return False
        return True

二、获取A2C的总损失并在更新后记录

A2C的总损失由策略损失、价值损失和熵损失加权计算而来,可在每次策略更新完成后触发的_on_policy_update方法中获取并计算:

  1. 通过self.locals字典直接取出A2C训练时生成的policy_loss、value_loss、entropy_loss分量;
  2. 按照A2C源码中的逻辑计算总损失(总损失 = 策略损失 + 价值损失 - 熵系数 * 熵损失);
  3. 完成记录或后续处理。

示例代码(承接上面的回调类):

def _on_policy_update(self) -> None:
        # 取出各损失分量并转换为数值
        policy_loss = self.locals['policy_loss'].item()
        value_loss = self.locals['value_loss'].item()
        entropy_loss = self.locals['entropy_loss'].item()
        
        # 计算总损失
        total_loss = policy_loss + value_loss - self.model.ent_coef * entropy_loss
        
        print(f"Policy update completed, total loss: {total_loss:.4f}")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 10:20:43