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

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字段存储动作分布的预测结果,可直接提取对应动作的预测概率/价值

代码示例

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 11:06:01