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

TensorFlow双图位移预测模型使用tf.Dataset训练时转置参数报错如何解决

错误原因

  • 核心是模型输入结构和数据集输出结构不匹配:你定义的模型输入是形状为(2,128,128,3)的单个张量,但你的tf.data pipeline输出的是([image, image2], dY),也就是长度为2的图像张量列表,而非拼接好的单张输入张量。
  • 两者结构错位导致模型收到的输入不符合预期,反向传播时transpose操作的维度校验失败,抛出了报错。

解决方案

两种可选方案任选其一即可:

方案一:修改模型输入,适配现有数据集输出

无需改动数据加载逻辑,仅调整大模型的输入定义即可,修改main函数中的对应代码:

# 替换原有的image_inputs定义、split、squeeze部分代码
first_image_input = keras.Input(shape=(128,128,3))
second_image_input = keras.Input(shape=(128,128,3))

image_outputs = [image_model(first_image_input), image_model(second_image_input)]
model = layers.Concatenate()(image_outputs)

# 后续层定义不变,最后修改模型输入为两个独立张量
final_model = keras.Model([first_image_input, second_image_input], out_layer)

方案二:修改数据pipeline,适配现有模型输入

如果要保持现有模型的单输入结构,就调整load_data_wrapper函数,在数据侧完成两张图像的拼接:

def load_data_wrapper(image_files):
    image, image2, dY = tf.py_function(load_data, [image_files], [tf.float32, tf.float32, tf.float32])
    # 沿axis=0堆叠两张图像,得到形状为(2,128,128,3)的输入张量
    combined_image = tf.stack([image, image2], axis=0)
    return (combined_image, dY)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 04:15:06