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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 01:45:16