基于TF Agents自定义PyEnvironment时Reward等形状不匹配报错问题
基于TF Agents官方DQN教程代码训练智能体,自定义PyEnvironment后,运行compute_avg_return函数时在policy.action(time_step)行触发错误。
报错代码
def compute_avg_return(environment, policy, num_episodes=10): total_return = 0.0 for _ in range(num_episodes): time_step = environment.reset() episode_return = 0.0 while not time_step.is_last(): action_step = policy.action(time_step) # <----- error on this line time_step = environment.step(action_step.action) episode_return += time_step.reward total_return += episode_return avg_return = total_return / num_episodes return avg_return.numpy()[0] compute_avg_return(eval_env, random_policy, num_eval_episodes)
报错信息
ValueError: Received a mix of batched and unbatched Tensors, or Tensors are not compatible with Specs. num_outer_dims: 1. Saw tensor_shapes: TimeStep( {'discount': TensorShape([1]), 'observation': TensorShape([1, 50, 30]), 'reward': TensorShape([1]), 'step_type': TensorShape([1])}) And spec_shapes: TimeStep( {'discount': TensorShape([]), 'observation': TensorShape([1, 50, 30]), 'reward': TensorShape([]), 'step_type': TensorShape([])})
从报错日志可见,observation形状符合要求,但discount、reward、step_type的形状与规格不匹配,需调整这些属性的形状。
核心问题是自定义PyEnvironment返回的discount、reward、step_type是带批量维度(shape [1])的张量,但TF Agents的Policy期望这些是无批量维度(shape [])的标量张量(与spec定义一致)。调整方法如下:
修改自定义Environment的
_reset和_step方法:
处理discount、reward、step_type时,去掉多余的批量维度。可以用tensorflow.squeeze()函数压缩张量:# 示例:在返回time_step前处理字段 reward = tf.squeeze(reward_tensor, axis=0) discount = tf.squeeze(discount_tensor, axis=0) step_type = tf.squeeze(step_type_tensor, axis=0)或者在创建这些张量时直接不添加批量维度,比如用
tf.constant(0.0)代替tf.constant([0.0])。对齐time_step_spec定义:
确保环境的time_step_spec中,discount、reward、step_type的spec为标量类型:time_step_spec = ts.TimeStep( step_type=tf.TensorSpec(shape=[], dtype=tf.int32), reward=tf.TensorSpec(shape=[], dtype=tf.float32), discount=tf.TensorSpec(shape=[], dtype=tf.float32), observation=tf.TensorSpec(shape=[1, 50, 30], dtype=tf.float32) )这里observation保留批量维度(符合当前场景),其余三个字段为标量shape []。
批量环境兼容提示:
若后续使用BatchedPyEnvironment批量环境,TF Agents会自动添加批量维度,此时无需手动处理;但当前单环境场景必须保持三个字段为标量。
内容的提问来源于stack exchange,提问作者laezZ_boi

