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

Keras多输入模型与生成器输入不匹配问题求助

问题解决:Keras双输入模型生成器输出结构不匹配

问题核心

你的生成器输出的是嵌套列表形式的numpy数组,而Keras双输入模型期望的是**((tf.float32张量, tf.float32张量), tf.float32标签张量)**的元组结构,两者的类型(numpy数组 vs TF张量)和结构(列表 vs 元组)都不匹配,导致训练报错。

解决步骤

1. 将numpy数组转换为TensorFlow张量

在生成器输出前,用tf.convert_to_tensor()把C++生成的numpy数组转换成TensorFlow的tf.float32类型张量。

2. 调整输出结构为元组嵌套

把生成器的输出从列表[[图像数组, 数值数组], 标签数组]改成元组((图像张量, 数值张量), 标签张量),贴合模型预期的输入结构。

修改后的生成器示例代码

假设你的原始生成器代码大致如下:

def data_generator():
    while True:
        # 通过Pybind11调用C++函数生成数据
        img_arr, num_arr, label_arr = generate_data_from_cpp()
        # 原始错误输出结构
        yield [[img_arr, num_arr], label_arr]

修改后:

import tensorflow as tf

def data_generator():
    while True:
        img_arr, num_arr, label_arr = generate_data_from_cpp()
        # 转换为tf.float32张量
        img_tensor = tf.convert_to_tensor(img_arr, dtype=tf.float32)
        num_tensor = tf.convert_to_tensor(num_arr, dtype=tf.float32)
        label_tensor = tf.convert_to_tensor(label_arr, dtype=tf.float32)
        # 输出符合要求的元组结构
        yield ((img_tensor, num_tensor), label_tensor)

额外验证建议

可以先单独调用生成器,打印输出的类型和结构,确认是否符合预期:

gen = data_generator()
sample_output = next(gen)
print("输出结构类型:", type(sample_output))
print("输入部分类型:", type(sample_output[0]))
print("图像张量类型:", type(sample_output[0][0]), sample_output[0][0].dtype)
print("数值张量类型:", type(sample_output[0][1]), sample_output[0][1].dtype)
print("标签张量类型:", type(sample_output[1]), sample_output[1].dtype)

运行后应看到所有张量都是tf.Tensor类型,dtype为float32,结构为((tf.Tensor, tf.Tensor), tf.Tensor)。


内容的提问来源于stack exchange,提问作者Balázs Bämer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 13:24:51