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

TensorFlow flow_from_directory训练灰度图分类报深度不匹配错误

灰度图CNN分类任务训练报错问题

问题背景

训练用于灰度图像6分类任务的CNN模型时,原本可在RGB图像数据集上正常运行的代码,切换到灰度图像数据集后触发报错,报错后IDE内核直接冻结。

报错信息

运行代码后抛出的错误提示为:

input depth must be evenly divisible by filter depth: 1 vs 3
报错截图如下:
CNN训练灰度图报错截图

相关代码

input_shape=(256, 256,1) # 最后一维设为1,适配灰度图单通道属性
target_size = (256, 256) # 图像缩放尺寸,供flow_from_directory接口调用

model_name='Test1'
model_filename = (model_name+'.hdf5') 

optimizer = Adam(learning_rate=1e-3)
loss=['categorical_crossentropy']
metrics = ['accuracy']

# 模型结构定义
model = Sequential()
model.add(Conv2D(32, (3, 3), input_shape=input_shape))
model.add(Activation('relu'))
model.add(MaxPooling2D(pool_size=(2, 2)))

model.add(Conv2D(64, (3, 3)))
model.add(Activation('relu'))
model.add(MaxPooling2D(pool_size=(2, 2)))

model.add(Flatten())
model.add(Dense(64))
model.add(Activation('relu'))
model.add(Dropout(0.5))

model.add(Dense(6)) # 输出维度对应6个分类
model.add(Activation('softmax'))

model.summary()

# 训练集数据增强配置
train_datagen = ImageDataGenerator(
        shear_range=0.2,
        zoom_range=0.2,
        horizontal_flip=True)

# 验证集数据生成器,不做数据增强
vaidation_datagen = ImageDataGenerator()

train_generator = train_datagen.flow_from_directory(
        train_path,  # 训练集图像所在文件夹路径
        target_size=target_size, 
        color_mode='grayscale', # 指定按灰度模式加载图像
        batch_size=batch_size,
        shuffle=True,
        class_mode='categorical',
        interpolation='nearest') 

validation_generator = vaidation_datagen.flow_from_directory(
        validation_path,  # 验证集图像所在文件夹路径
        target_size=target_size,
        color_mode='grayscale', # 指定按灰度模式加载图像
        batch_size=batch_size,
        shuffle=True,
        class_mode='categorical',
        interpolation='nearest')

model.compile(optimizer, loss , metrics)
# 模型 checkpoint 回调配置
model_checkpoint = tf.keras.callbacks.ModelCheckpoint((model_path+model_filename), monitor='loss',verbose=1, save_best_only=True)
model.summary()

# 启动模型训练
history = model.fit(
     train_generator,
     steps_per_epoch = num_of_train_img_raw//batch_size,
     epochs = epochs, 
     validation_data = validation_generator,
     validation_steps = num_of_val_img_raw//batch_size,
     callbacks=[model_checkpoint],
     use_multiprocessing = False)

报错原因

报错核心是输入张量通道数和卷积核要求的通道数不匹配:代码虽然已经将模型输入通道设为1、数据生成器也指定了灰度加载模式,但如果运行前没有清空之前训练RGB模型时残留的权重、计算图,或者代码隐式加载了之前在RGB数据集上训练保存的3通道权重,第一层卷积的卷积核通道数为3,和当前1通道的输入维度不匹配,就会抛出该错误,内核冻结通常是显存被残留的旧计算图占满导致的。

解决方法

  • 先完全重启Python内核,清空显存中残留的所有旧模型、旧变量,重新运行代码时不要加载之前RGB训练阶段保存的.hdf5权重文件
  • 如果需要复用RGB数据集上预训练的权重做迁移学习,不要使用单通道输入:将数据生成器的color_mode改为rgb,input_shape改回(256,256,3),加载灰度图后将单通道数据复制3份拼接成3通道伪RGB张量再输入模型
  • 检查代码中是否有遗漏的预训练权重加载逻辑,比如Conv2D层是否默认开启了预训练权重加载(如ImageNet预训练权重,默认适配3通道输入),如有则关闭该参数
  • 修正代码里的拼写错误:验证集生成器变量名vaidation_datagen正确拼写应为validation_datagen,避免后续调用触发变量不存在的错误

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 22:45:39