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

基于Keras数据集构建双通道图像分类ANN的报错排查与优化

问题描述

我正尝试设计一个双通道ANN,以两种类型的图像作为输入,将其分类为5个类别。数据集结构如下:

spc_trn_dir/
...class_1/
......image_1.jpg
......image_2.jpg
...class_2/
......image_11.jpg
......image_12.jpg
...class_3/
......image_21.jpg
......image_22.jpg
...class_4/
......image_31.jpg
......image_32.jpg
...class_5/
......image_41.jpg
......image_42.jpg

scl_trn_dir/
...class_1/
......image_51.jpg
......image_52.jpg
...class_2/
......image_61.jpg
......image_62.jpg
...class_3/
......image_71.jpg
......image_72.jpg
...class_4/
......image_81.jpg
......image_82.jpg
...class_5/
......image_91.jpg
......image_92.jpg

数据集包含8700张150×150的图像,无法用NumPy数组加载(会耗尽内存),因此采用Keras数据集方案。

已设计的网络结构如下:

Model: "model_9"
__________________________________________________________________________________________________
 Layer (type)                   Output Shape         Param #     Connected to                     
==================================================================================================
 img_input_1 (InputLayer)       [(None, 150, 150, 3)]  0           []                               
                                                                                                    
 img_input_2 (InputLayer)       [(None, 150, 150, 3)]  0           []                               
                                                                                                    
 model_2 (Functional)           (None, 1024)         2730040     ['img_input_1[0][0]']             
                                                                                                    
 model_3 (Functional)           (None, 1024)         2730040     ['img_input_2[0][0]']             
                                                                                                    
 concatenate_10 (Concatenate)   (None, 2048)         0           ['model_2[11][][0]',               
                                                                  'model_3[11][0]']               
                                                                                                    
 dense_9 (Dense)                (None, 5)            10245       ['concatenate_10[0][0]']          
                                                                                                    
==================================================================================================
Total params: 5,470,325
Trainable params: 5,453,749
Non-trainable params: 16,576
__________________________________________________________________________________________________

构建数据集的代码:

image_size = (150, 150)
batch_size = 128

spc_train_ds = tf.keras.preprocessing.image_dataset_from_directory(
    spc_trn_dir,
    image_size=image_size,
    batch_size=batch_size,
)

scl_train_ds = tf.keras.preprocessing.image_dataset_from_directory(
    scl_trn_dir,
    image_size=image_size,
    batch_size=batch_size,
)

train_ds = tf.data.Dataset.zip((spc_train_ds,scl_train_ds))
val_ds = tf.data.Dataset.zip((spc_val_ds,scl_val_ds))

验证数据集输出形状:

(spc_img, spc_lbl),( scl_img,  scl_lbl)= next(iter(train_ds))
print(f'shapes: image batch: {spc_img.shape} , labels: {spc_lbl.shape}')
print(f'shapes: image batch: {scl_img.shape} , labels: {scl_lbl.shape}')

输出:

shapes: image batch: (128, 150, 150, 3) , labels: (128,)
shapes: image batch: (128, 150, 150, 3) , labels: (128,)

执行multi_modal_model.fit(train_ds, epochs=1, validation_data=val_ds)时,出现输入形状不兼容错误:

WARNING:tensorflow:Model was constructed with shape (None, 150, 150, 3) for input KerasTensor(type_spec=TensorSpec(shape=(None, 150, 150, 3), dtype=tf.float32, name='img_input_2'), name='img_input_2', description="created by layer 'img_input_2'"), but it was called on an input with incompatible shape (None,).
WARNING:tensorflow:Model was constructed with shape (None, 150, 150, 3) for input KerasTensor(type_spec=TensorSpec(shape=(None, 150, 150, 3), dtype=tf.float32, name='input_6'), name='input_6', description="created by layer 'input_6'"), but it was called on an input with incompatible shape (None,).

ValueError: Exception encountered when calling layer "model_3" (type Functional).

Input 0 of layer "conv2d_12" is incompatible with the layer: expected min_ndim=4, found ndim=1. Full shape received: (None,)

Call arguments received by layer "model_3" (type Functional):
  • inputs=tf.Tensor(shape=(None,), dtype=float32)
  • training=True
  • mask=None

已知模型误将第一个通道的标签作为第二个通道的输入,需要修正该问题,并获得更优的双通道数据集构建方案。


错误原因

tf.data.Dataset.zip合并两个数据集后,得到的结构是((spc_img, spc_lbl), (scl_img, scl_lbl))。而Keras的fit方法会自动将输入结构拆分为模型输入和标签,误把spc_lbl当成了模型的第二个输入张量,导致卷积层收到形状为(None,)的标签数据,引发维度不匹配错误。


解决方法

需要重新整理数据集结构,将输入转换为两个图像张量组成的元组,标签取其中一组(两类图像对应同一类别,标签一致)。具体实现如下:

# 定义数据映射函数,整理输入输出结构
def restructure_data(inputs):
    (spc_img, spc_lbl), (scl_img, scl_lbl) = inputs
    # 两类图像标签一致,取任意一组作为模型标签
    return (spc_img, scl_img), spc_lbl

# 应用映射函数处理训练集和验证集
train_ds = tf.data.Dataset.zip((spc_train_ds, scl_train_ds)).map(restructure_data)
val_ds = tf.data.Dataset.zip((spc_val_ds, scl_val_ds)).map(restructure_data)

验证处理后的数据集结构:

(inputs, labels) = next(iter(train_ds))
spc_img, scl_img = inputs
print(f'spc_img shape: {spc_img.shape}, scl_img shape: {scl_img.shape}, labels shape: {labels.shape}')

预期输出:

spc_img shape: (128, 150, 150, 3), scl_img shape: (128, 150, 150, 3), labels shape: (128,)

此时数据集结构完全匹配模型的输入要求,可以正常执行fit训练。


优化的双通道数据集构建方案

针对大规模图像数据集,可从数据增强、性能优化、数据校验三个维度优化构建流程:

1. 加入数据增强

针对两类图像分别添加随机增强操作,提升模型泛化能力:

# 定义数据增强流水线
data_augmentation = tf.keras.Sequential([
    tf.keras.layers.RandomFlip("horizontal"),
    tf.keras.layers.RandomRotation(0.1),
    tf.keras.layers.RandomZoom(0.1),
])

2. 标准化与性能优化

将像素值标准化到[0,1]区间,并加入缓存、预取操作,避免IO瓶颈:

def load_dataset(directory, augment=False):
    ds = tf.keras.preprocessing.image_dataset_from_directory(
        directory,
        image_size=(150, 150),
        batch_size=128,
        label_mode='int'
    )
    # 像素值标准化
    ds = ds.map(lambda x, y: (tf.cast(x, tf.float32) / 255.0, y))
    # 应用数据增强(仅训练集)
    if augment:
        ds = ds.map(lambda x, y: (data_augmentation(x, training=True), y))
    # 缓存数据+预取,提升训练速度
    ds = ds.cache().prefetch(tf.data.AUTOTUNE)
    return ds

# 加载训练集与验证集
spc_train_ds = load_dataset(spc_trn_dir, augment=True)
scl_train_ds = load_dataset(scl_trn_dir, augment=True)
spc_val_ds = load_dataset(spc_val_dir)
scl_val_ds = load_dataset(scl_val_dir)

3. 数据对齐校验

加入标签一致性校验,避免两类图像的样本对齐错误:

def merge_and_validate(inputs):
    (spc_img, spc_lbl), (scl_img, scl_lbl) = inputs
    # 校验两类图像的标签是否一致
    tf.debugging.assert_equal(spc_lbl, scl_lbl, message="两类图像标签不匹配,样本对齐错误!")
    return (spc_img, scl_img), spc_lbl

train_ds = tf.data.Dataset.zip((spc_train_ds, scl_train_ds)).map(merge_and_validate)
val_ds = tf.data.Dataset.zip((spc_val_ds, scl_val_ds)).map(merge_and_validate)

优化后的优势

  • 数据增强:提升模型对不同视角、尺度图像的识别能力
  • 标准化:加快模型收敛速度
  • 缓存+预取:让数据加载与模型训练并行,大幅提升训练效率
  • 标签校验:提前排查数据对齐错误,避免无效训练

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 17:05:28