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

调用TensorFlow DQN Agent collect_policy时出现维度不匹配错误

问题排查:TensorFlow DQN Agent collect_policy动作调用报错

问题描述

在Simulink环境中使用TensorFlow DQN Agent,调用agent.collect_policy.action(time_step)时触发InvalidArgumentError,提示'then'和'else'必须尺寸相同,但实际为[1] vs [];调用agent.policy.action(time_step)则正常运行,已确认TimeStep与TimeStepSpec匹配。

报错信息

tensorflow.python.framework.errors_impl.InvalidArgumentError: {{function_node __wrapped__Select_device_/job:localhost/replica:0/task:0/device:CPU:0}} 'then' and 'else' must have the same size.  but received: [1] vs. [] [Op:Select] name:

代码片段

规格定义

discount = 0.95
reward = 0.0
optimizer = tf.keras.optimizers.Adam(learning_rate=learning_rate)

time_step_spec = TimeStep(step_type = tensor_spec.BoundedTensorSpec(shape=(1,), dtype=tf.int32, minimum=0, maximum=2),
                        reward = tensor_spec.TensorSpec(shape=(1,), dtype=tf.float32),
                        discount = tensor_spec.TensorSpec(shape=(1,), dtype=tf.float32),
                        observation =  tensor_spec.TensorSpec(shape=(1,amountMachines), dtype=tf.int32)
                        )

num_possible_actions = 729
action_spec = tensor_spec.BoundedTensorSpec(
    shape=(), dtype=tf.int32, minimum=0, maximum=num_possible_actions - 1)

agent = dqn_agent.DqnAgent(
    time_step_spec,
    action_spec,
    q_network=model,
    optimizer=optimizer,
    epsilon_greedy= 1.0,
    td_errors_loss_fn=common.element_wise_squared_loss,
    train_step_counter=train_step_counter)
agent.initialize()

调用代码

current_state = get_states() # 获取形如[4,4,4,4,4,6]的np.array
current_state_batch = tf.expand_dims( tf.convert_to_tensor(current_state, dtype=tf.int32), axis=0)

time_step = TimeStep(step_type=tf.convert_to_tensor([step_type], dtype=tf.int32),
                    reward=tf.convert_to_tensor([reward], dtype=tf.float32),
                    discount=tf.convert_to_tensor([discount], dtype=tf.float32),
                    observation= current_state_batch)

action_step = agent.collect_policy.action(time_step)

完整报错栈

Traceback (most recent call last):
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\runpy.py", line 198, in _run_module_as_main
    return _run_code(code, main_globals, None,
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\runpy.py", line 88,  in _run_code
    exec(code, run_globals)
  File "c:\Users\STestUser\.vscode\extensions\ms-python.python-2023.20.0\pythonFiles\lib\python\debugpy\adapter/../..\debugpy\launcher/../..\debugpy\__main__.py", line 39, in <module>
    cli.main()
  File "c:\Users\STestUser\.vscode\extensions\ms-python.python-2023.20.0\pythonFiles\lib\python\debugpy\adapter/../..\debugpy\launcher/../..\debugpy\..\debugpy\server\cli.py", line 430, in main
    run()
  File "c:\Users\STestUser\.vscode\extensions\ms-python.python-2023.20.0\pythonFiles\lib\python\debugpy\adapter/../..\debugpy\launcher/../..\debugpy\..\debugpy\server\cli.py", line 284, in run_file
    runpy.run_path(target, run_name="__main__")
  File "c:\Users\STestUser\.vscode\extensions\ms-python.python-2023.20.0\pythonFiles\lib\python\debugpy\_vendored\pydevd\_pydevd_bundle\pydevd_runpy.py", line 321, in run_path
    return _run_module_code(code, init_globals, run_name,
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "c:\Users\STestUser\.vscode\extensions\ms-python.python-2023.20.0\pythonFiles\lib\python\debugpy\_vendored\pydevd\_pydevd_bundle\pydevd_runpy.py", line 135, in _run_module_code
    _run_code(code, mod_globals, init_globals,
  File "c:\Users\STestUser\.vscode\extensions\ms-python.python-2023.20.0\pythonFiles\lib\python\debugpy\_vendored\pydevd\_pydevd_bundle\pydevd_runpy.py", line 124, in _run_code
    exec(code, run_globals)
  File "d:\Hochschule\Master\Masterarbeit\energy-efficiency-optimation\RL-Modell\OP10_QLearning.py", line 449, in <module>
    action_step = agent.collect_policy.action(time_step = time_step_t)        
                  ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tf_agents\policies\tf_policy.py", line 333, in action
    step = action_fn(time_step=time_step, policy_state=policy_state, seed=seed)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tf_agents\utils\common.py", line 193, in with_check_resource_vars
    return fn(*fn_args, **fn_kwargs)
           ^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tf_agents\policies\epsilon_greedy_policy.py", line 141, in _action
    action = tf.nest.map_structure(
             ^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tensorflow\python\util\nest.py", line 629, in map_structure
    return nest_util.map_structure(
           ^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tensorflow\python\util\nest_util.py", line 1168, in map_structure
    return _tf_core_map_structure(func, *structure, **kwargs)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tensorflow\python\util\nest_util.py", line 1208, in _tf_core_map_structure
    [func(*x) for x in entries],
    ^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tensorflow\python\util\nest_util.py", line 1208, in <listcomp>
    [func(*x) for x in entries],
     ^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tf_agents\policies\epsilon_greedy_policy.py", line 142, in <lambda>
    lambda g, r: tf.compat.v1.where(cond, g, r),
                 ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tensorflow\python\util\traceback_utils.py", line 153, in error_handler
    raise e.with_traceback(filtered_tb) from None
  File "C:\Users\STestUser\AppData\Local\anaconda3\Lib\site-packages\tensorflow\python\framework\ops.py", line 5888, in raise_from_not_ok_status
    raise core._status_to_exception(e) from None  # pylint: disable=protected-access
    ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
tensorflow.python.framework.errors_impl.InvalidArgumentError: {{function_node __wrapped__Select_device_/job:localhost/replica:0/task:0/device:CPU:0}} 'then' and 'else' must have the same size.  but received: [1] vs. [] [Op:Select] name:

问题分析与解决方法

从报错栈可定位问题出在epsilon_greedy_policy.py的tf.compat.v1.where调用:epsilon贪婪策略需要在贪婪动作和随机动作间做选择,但两者张量形状不匹配:

  • 贪婪动作由DQN网络输出,因输入是batch维度(shape=(1,))的观测,输出动作形状为[1]
  • 随机动作基于action_spec生成,而你的action_spec定义为shape=()(标量),所以随机动作形状是[]

两者形状不一致导致tf.where无法执行选择操作,而agent.policy是纯贪婪策略,无需生成随机动作,因此运行正常。

解决步骤:

  1. 统一动作形状:将action_spec的形状改为与贪婪动作一致的(1,):

    action_spec = tensor_spec.BoundedTensorSpec(
        shape=(1,), dtype=tf.int32, minimum=0, maximum=num_possible_actions - 1)
    

    这样随机动作会生成[1]形状的张量,与贪婪动作匹配。

  2. 检查DQN网络输出:确保Q网络输出的动作维度与修改后的action_spec一致,避免后续训练或推理出现形状不兼容问题。

  3. 验证TimeStep一致性:确认所有TimeStep字段的batch维度保持统一,避免因部分字段是标量、部分是batch张量导致的隐性形状问题。

内容的提问来源于stack exchange,提问作者Simon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 09:49:51