训练TF Agent调用PyDriver.run时遇批次维度ValueError求助
解决TF Agent中PyDriver.run触发的ValueError问题
错误根源
- 创建
collect_driver时使用的是原生Python环境(env = myEnv()),但传入的初始time_step来自TFPyEnvironment实例(train_env.reset()),后者返回的time_step带有batch维度(shape=(1, ...)) PyDriver配合原生Python环境工作时,要求输入的time_step是无batch维度的单步数据,而EpsilonGreedyPolicy不支持多batch维度的输入,因此触发报错。
解决方案
方案1:修正PyDriver的time_step来源(适配原生环境)
将初始time_step的获取改为从原生环境env获取,同时显式配置Policy处理单步数据:
# 替换原time_step的获取方式 time_step = env.reset() # 原生环境返回无batch维度的time_step # 创建PyDriver时显式设置batch_time_steps=False collect_driver = py_driver.PyDriver( env, py_tf_eager_policy.PyTFEagerPolicy( agent.collect_policy, use_tf_function=True, batch_time_steps=False), [rb_observer], max_steps=collect_steps_per_iteration)
方案2:改用TFDriver(推荐,贴合TF Agent工作流)
TFDriver专门为TFPyEnvironment设计,天然支持批量数据,和你的agent训练环境完全匹配:
# 导入TFDriver from tf_agents.drivers import tf_driver # 替换PyDriver为TFDriver collect_driver = tf_driver.TFDriver( train_env, agent.collect_policy, observers=[rb_observer], max_steps=collect_steps_per_iteration) # 初始time_step仍从train_env获取(带batch维度,符合TFDriver要求) time_step = train_env.reset()
说明
- 方案1适合需要保留原生Python环境的场景,确保输入输出维度和PyDriver的要求一致。
- 方案2是更推荐的方式,因为你的agent是基于
TFPyEnvironment创建的,使用TFDriver能更好地融入TF Agent的批量训练工作流,避免维度不匹配问题。
内容的提问来源于stack exchange,提问作者변상진
相关产品推荐
相关产品推荐

