TF-Agents Deep Q Learning:如何提取状态/动作对的预测值?
TF Agents 读取SavedModel策略动作关联预测值的方法
你通过SavedModelPyTFEagerPolicy加载的TF Agents策略原生支持提取动作关联的预测值,无需额外调用第三方函数,操作逻辑如下:
核心原理
TF Agents的策略类调用action()方法时,默认返回PolicyStep命名元组,包含三个字段:
action:你已经能正常获取的预测动作结果state:循环策略(RNN结构)的隐藏状态,非循环策略返回空占位值info:你需要的预测值默认存储在该字段下,不同策略类型对应字段不同:- DQN等值迭代类策略:
info下的q_values字段存储所有动作对应的预测Q值,按动作索引取值即可得到当前动作的预测Q值 - PPO等策略梯度类策略:
info下的distribution字段存储动作分布的预测结果,可直接提取对应动作的预测概率/价值
- DQN等值迭代类策略:
代码示例
from tf_agents.policies import saved_model_py_tfeager_policy # 加载已训练好的SavedModel策略(和你现有加载逻辑一致即可) policy = saved_model_py_tfeager_policy.SavedModelPyTFEagerPolicy( "your_policy_saved_path", load_specs_from_pbtxt=True ) # 构造测试用的TimeStep格式观测值,和你之前提取动作的入参规格一致 test_time_step = 你的测试观测数据 # 调用action方法获取完整返回结果 policy_step = policy.action(time_step=test_time_step) # 提取动作关联预测值示例 ## DQN类策略取当前动作的预测Q值 q_values = policy_step.info.q_values target_action_q = q_values.numpy()[0][policy_step.action.numpy()] # 按实际batch维度调整索引 ## 策略梯度类策略取当前动作的预测概率 action_dist = policy_step.info.distribution target_action_prob = action_dist.prob(policy_step.action).numpy()
常见问题排查
如果提取info下字段时报错,通常是导出SavedModel时没有同步导出info字段:
- 回到训练脚本,导出模型时显式指定要保留policy info字段,重新导出后加载即可正常读取
- 可调用
policy.info_spec()查看当前加载的策略支持返回的所有预测信息字段,匹配你需要的内容提取即可
内容的提问来源于stack exchange,提问作者Setjmp
相关产品推荐
相关产品推荐

