如何使用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
相关产品推荐
相关产品推荐

