StableBaselines3:如何在env.reset()后创建自定义回调函数
实现环境重置后触发的自定义逻辑
Stable Baselines3(SB3)官方提供的回调确实没有专门对应「环境重置后」的触发钩子,但有两种简单的方式实现需求:
方法一:直接修改自定义环境的reset方法
这是最直接可靠的方式,所有环境重置操作(不管是SB3训练时自动调用,还是手动触发)都会走这个方法。只需要在自定义环境类里重写reset函数,先完成原生重置逻辑,再追加你的自定义操作:
import gymnasium as gym from gymnasium import spaces class CustomEnv(gym.Env): def __init__(self): super().__init__() # 你的环境初始化逻辑 self.observation_space = spaces.Box(low=0, high=1, shape=(4,)) self.action_space = spaces.Discrete(2) def reset(self, seed=None, options=None): # 先执行原生重置逻辑 obs, info = super().reset(seed=seed, options=options) # 这里写你要在重置后执行的代码 print("环境已重置,执行自定义逻辑") # 比如记录重置次数、初始化某些环境变量等 return obs, info def step(self, action): # 你的step逻辑 obs = self.observation_space.sample() reward = 0.0 terminated = False truncated = False info = {} return obs, reward, terminated, truncated, info
方法二:利用SB3的EpisodeStart回调
如果不想修改环境代码,可以用SB3的BaseCallback,重写_on_episode_start方法——SB3在每个新episode开始前会自动调用环境重置,所以这个方法的执行时机刚好对应「重置后、新episode第一步前」,基本满足需求:
from stable_baselines3.common.callbacks import BaseCallback class ResetCallback(BaseCallback): def _on_episode_start(self, locals_, globals_): # 这里写重置后要执行的逻辑 print("新episode开始(环境已重置),执行自定义逻辑") # 可以通过locals_获取环境实例:env = locals_["env"] env = locals_["env"] # 比如操作环境的内部变量 # env.some_variable = 0 # 训练时加载回调 from stable_baselines3 import PPO env = CustomEnv() model = PPO("MlpPolicy", env, verbose=1) callback = ResetCallback() model.learn(total_timesteps=10000, callback=callback)
注意:如果你的场景中存在**手动调用env.reset()**的情况(比如训练中途手动重置环境),方法二的_on_episode_start不会触发,这种情况下优先用方法一。
内容的提问来源于stack exchange,提问作者tirilazat
相关产品推荐
相关产品推荐

