[错误]: 使用 EvalCallback 进行渲染时,不会渲染初始状态或最终状态

作者: kylesayrs创建于 2023年9月23日更新于 2026年5月26日
标签bug

To Reproduce

from typing import Optional import numpy as np from stable_baselines3.common.env_util import make_vec_env from stable_baselines3 import DQN from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.callbacks import EvalCallback from stable_baselines3.common.envs import BitFlippingEnv

A user should not be required to subclass their environment to debug the EvalCallback

class ClearerRenderEnv(BitFlippingEnv): def step(self, *args, **kwargs): pre_state = self.state.copy() rets = super().step(*args, **kwargs) print(f"step : {pre_state} -> {self.state} | {self.desired_goal}")

    return rets

def reset(self, *args, **kwargs):
    if not hasattr(self, "state"):
        return super().reset(*args, **kwargs)
    
    pre_state = self.state.copy()
    rets = super().reset(*args, **kwargs)
    print(f"reset   : {pre_state} -> {self.state} | {self.desired_goal} ")

    return rets

def render(self) -> Optional[np.ndarray]:
    if self.render_mode == "rgb_array":
        return self.state.copy()
    print(f"rendered:          {self.state} | {self.desired_goal}")

if name == "main": eval_callback = EvalCallback( Monitor(ClearerRenderEnv(n_bits=2)), n_eval_episodes=1, eval_freq=10_000, render=True, )

environment = make_vec_env(
    BitFlippingEnv,
    env_kwargs={"n_bits": 2},
    n_envs=1
)

model = DQN(
    "MultiInputPolicy",
    environment,
    learning_starts=0,
    learning_rate=0.1,
    verbose=2
)

model.learn(
    total_timesteps=10_000,
    log_interval=None,
    callback=eval_callback,
    progress_bar=True,
)

Relevant log output / Error message

Notice the first state [0, 0] is never rendered and neither is the last step, only the reset frame (in this case, also [0, 0])

step    : [0 0] -> [0 1] | [1 1]
rendered:          [0 1] | [1 1]
step    : [0 1] -> [1 1] | [1 1]
reset   : [1 1] -> [0 0] | [1 1] 
rendered:          [0 0] | [1 1]
Eval num_timesteps=10000, episode_reward=-1.00 +/- 0.00
Episode length: 2.00 +/-
…

内容来源: DLR-RM/stable-baselines3