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

如何使用Keras Callback记录强化学习模型测试时每步返回的观测值

问题修复与实现方案

原代码的核心问题

  • 语法缩进错误:类内部的__init__、on_step_end方法没有按Python规范缩进,直接运行会触发语法错误
  • 存储结构初始化缺失:self.observations初始化为空字典,直接对不存在的step键执行append操作会触发KeyError
  • 没有获取真实观测值:硬编码的['observation']是固定字符串列表,无法存储模型测试过程中返回的实际观测
  • 无episode区分逻辑:多episode测试时,不同episode的同step数据会互相覆盖,无法对应到各自的episode序列

可运行的实现代码

如果你使用的是keras-rl库的DQN接口,参考实现如下:

from keras.callbacks import Callback

class StepLogger(Callback):
    def __init__(self):
        # 存储结构:外层键为episode序号,值为该episode所有step的观测列表
        self.episode_observations = {}
        self.current_episode = 0

    def on_episode_begin(self, episode, logs=None):
        # 新episode启动时,初始化该episode的观测存储列表
        self.current_episode = episode
        self.episode_observations[episode] = []

    def on_step_end(self, step, logs=None):
        # 从logs中读取当前step的实际观测,keras-rl默认会将观测写入logs字典
        current_observation = logs.get('observation')
        # 追加到当前episode的观测序列中
        self.episode_observations[self.current_episode].append(current_observation)

# 初始化回调实例,测试完成后直接从该实例读取存储的观测
step_logger = StepLogger()
dqn.test(env, nb_episodes=2, visualize=False, callbacks=[step_logger])

# 测试完成后验证数据
for episode_id, obs_sequence in step_logger.episode_observations.items():
    print(f"第{episode_id+1}个episode总步数:{len(obs_sequence)}")
    print(f"首步观测值:{obs_sequence[0]}")

补充说明

  • 如果你使用的keras-rl版本没有默认将观测写入logs,可以自行修改DQN的test方法,在每步执行完env.step后将观测写入传给回调的logs字典即可
  • 如需持久化存储观测数据,测试完成后可以用pickle将episode_observations序列化保存到本地文件

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 16:18:02