在Stable-Baselines3自定义回调中访问A2C指标与环境状态
问题解答
一、获取环境的truncated/terminated状态,实现截断终止训练与记录episode时长
完全可以在自定义回调的_on_step方法中获取这两个状态,具体实现如下:
- 先在回调类中维护一个变量,用于记录当前episode的步数;
- 通过
self.locals['infos']获取每个环境的信息字典,其中包含terminated和truncated布尔字段,对应环境的终止/截断状态; - 当检测到环境进入
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方法中获取并计算:
- 通过
self.locals字典直接取出A2C训练时生成的policy_loss、value_loss、entropy_loss分量; - 按照A2C源码中的逻辑计算总损失(总损失 = 策略损失 + 价值损失 - 熵系数 * 熵损失);
- 完成记录或后续处理。
示例代码(承接上面的回调类):
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
相关产品推荐
相关产品推荐

