使用ImageDataGenerator准备语义分割数据时报ValueError如何解决?
错误原因
flow_from_directory方法默认返回的每批次数据为(样本张量, 分类标签)二元组,其中分类标签是根据目录结构生成的类别编号,并非你需要的掩码数据。你对图像生成器和掩码生成器执行zip操作后,每批次返回的内容为((输入图像, 无用分类标签1), (掩码图像, 无用分类标签2)),模型会把前两个元素都识别为输入张量,因此抛出“收到2个输入张量但仅需要1个”的报错。- 额外潜在问题:掩码默认会以RGB三通道模式读取,不符合语义分割任务通常需要单通道掩码计算损失的要求;同时随机增强的种子如果没有全局统一,可能导致输入图像和对应掩码的增强变换不一致,出现标注错位。
修复方案
调整data_aug函数的参数,取消自动生成的分类标签输出,再对生成器做包装,返回符合(输入, 标签)要求的批次格式,修改后代码如下:
from tensorflow.keras.preprocessing.image import ImageDataGenerator def data_aug(batch_size=32, seed=42): datagen = ImageDataGenerator(rotation_range=10, validation_split=0.2) # 训练输入图像生成器:仅返回图像张量,丢弃自动生成的分类标签 X_train_augmented = datagen.flow_from_directory( directory='../input/train/fg_image', target_size=(256, 256), shuffle=True, batch_size=batch_size, class_mode=None, # 关键修改:不返回目录对应的分类标签 seed=seed # 统一随机种子,保证图像和掩码增强逻辑一致 ) # 训练掩码生成器 Y_train_augmented = datagen.flow_from_directory( directory='../input/train/gt_mask', target_size=(256, 256), shuffle=True, batch_size=batch_size, class_mode=None, color_mode='grayscale', # 单通道掩码,若你的掩码是三通道可修改为'rgb' seed=seed ) # 验证集输入图像生成器 X_val_augmented = datagen.flow_from_directory( directory='../input/validation/fg_image', target_size=(256, 256), shuffle=True, batch_size=batch_size, class_mode=None, seed=seed ) # 验证集掩码生成器 Y_val_augmented = datagen.flow_from_directory( directory='../input/validation/gt_mask', target_size=(256, 256), shuffle=True, batch_size=batch_size, class_mode=None, color_mode='grayscale', seed=seed ) # 包装生成器,输出符合训练要求的(输入, 标签)二元组格式 def pair_generator(img_gen, mask_gen): while True: img = img_gen.next() mask = mask_gen.next() # 可在此处添加掩码归一化、标签重编码等自定义预处理 yield (img, mask) train_generator = pair_generator(X_train_augmented, Y_train_augmented) val_generator = pair_generator(X_val_augmented, Y_val_augmented) # 返回生成器和样本数,用于计算训练步长 return train_generator, val_generator, X_train_augmented.samples, X_val_augmented.samples
训练代码同步修改:
train_generator, val_generator, train_samples, val_samples = data_aug(batch_size=32) # 新版Keras已废弃fit_generator,直接使用fit方法即可支持生成器输入 model.fit( train_generator, steps_per_epoch=train_samples // 32, # 按实际样本数计算步长,避免数据重复/缺失 epochs=50, validation_data=val_generator, validation_steps=val_samples // 32 )
内容的提问来源于stack exchange,提问作者MSI
相关产品推荐
相关产品推荐

