如何加载tf Agent中由triggers.PolicySavedModelTrigger保存的策略?
评估策略选择
评估就选greedy_policy文件夹下的策略就行,四个文件夹的用途差异如下:
checkpoints:存的是训练全流程状态快照,包含智能体权重、回放缓存、训练步数等完整信息,只用来做断点续训,不适合单独拉出来做策略评估collect_policy:是训练阶段收集数据用的策略,会主动加探索噪声,动作输出有随机波动,不符合评估要测最优效果的需求greedy_policy:训练出来的确定性贪心策略,不带探索噪声,会输出当前策略下的最优动作,是标准评估场景的首选policy:原始的随机策略实例,会基于概率分布采样动作,只有你需要评估随机策略的效果时才需要选这个
具体加载方法
直接用TensorFlow的SavedModel接口加载就行,操作代码示例如下:
- 先导入依赖
import tensorflow as tf
- 加载策略
# 把下面的路径替换成你本地greedy_policy文件夹的实际存储路径 policy_path = "./policies/greedy_policy" loaded_policy = tf.saved_model.load(policy_path)
- 评估时调用示例
# 拿到当前环境的时间步输入 current_time_step = your_env.current_time_step() # 调用加载好的策略得到动作输出 action_step = loaded_policy.action(current_time_step) # 转成numpy格式的动作值即可用于后续评估逻辑 eval_action = action_step.action.numpy()
内容的提问来源于stack exchange,提问作者Quan Vuong
相关产品推荐
相关产品推荐

