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

使用CustomDataGenerator与Keras模型时出现输入数量错误

错误原因分析

你的报错核心是模型期望2个输入张量,但数据生成器返回的输入格式不符合要求,具体问题如下:

  • 数据类型不规范:生成器返回的batch_images和batch_coordinates是Python列表,而非numpy数组/TensorFlow张量,Keras无法正确解析这种嵌套结构,误将列表元素拆分为多个独立输入。
  • 图像尺寸不匹配:InceptionV3要求输入图像固定为(299, 299, 3)尺寸,但你直接读取原图未做统一resize,会导致输入形状不一致。
  • 坐标形状未对齐:批量坐标需要是(batch_size, 2)的二维数组,而你返回的是一维列表,无法匹配模型input_coordinates层的输入形状。
修正后的代码

1. 修复CustomDataGenerator类

from tensorflow.keras.applications import InceptionV3
from tensorflow.keras.layers import Dense, Flatten, Input, concatenate
from tensorflow.keras.models import Model
from tensorflow.keras.utils import Sequence, to_categorical
import numpy as np
import cv2

class CustomDataGenerator(Sequence):
    def __init__(self, image_filenames, coordinates, labels, batch_size, img_size=(299, 299)):
        self.image_filenames = image_filenames
        self.coordinates = np.array(coordinates)  # 提前转为numpy数组
        self.labels = np.array(labels)
        self.batch_size = batch_size
        self.img_size = img_size

    def __len__(self):
        return len(self.image_filenames) // self.batch_size

    def __getitem__(self, index):
        start_idx = index * self.batch_size
        end_idx = (index + 1) * self.batch_size
        
        batch_filenames = self.image_filenames[start_idx:end_idx]
        batch_coords = self.coordinates[start_idx:end_idx]
        batch_labels = self.labels[start_idx:end_idx]

        # 统一处理图像尺寸并转为numpy数组
        batch_images = []
        for filename in batch_filenames:
            img = cv2.imread(filename)
            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
            img = cv2.resize(img, self.img_size)  # 匹配InceptionV3输入尺寸
            batch_images.append(img)
        batch_images = np.array(batch_images)

        # 返回标准格式:(输入张量列表, 标签数组)
        return [batch_images, batch_coords], batch_labels

2. 处理标签格式(针对categorical_crossentropy)

如果你的原始标签是整数形式(如0、1、2),需要转为独热编码:

# 假设train_labels、val_labels是整数列表
train_labels = to_categorical(train_labels, num_classes=3)
val_labels = to_categorical(val_labels, num_classes=3)

3. 模型定义与训练部分不变

原模型的结构定义无需修改,直接使用修正后的生成器执行训练即可。

关键修正点说明
  • 提前将坐标和标签转为numpy数组,避免批量处理时的类型混乱。
  • 强制统一图像尺寸,确保符合InceptionV3的输入要求。
  • 将批量图像转为numpy数组,让Keras正确识别为单个输入张量,而非零散的图像元素。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 21:22:57