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

如何将ImageDataGenerator传入U-net实现多类别分割及相关疑问

多类别分割中ImageDataGenerator与U-Net配合问题解答

问题描述

尝试使用ImageDataGenerator配合segmentation_models的U-Net实现多类别分割,编写代码如下:

data_generator  = ImageDataGenerator(
                    rescale = 1./255.
)

train_dataset_images = data_generator.flow_from_directory(
                            directory=image_directory,
                            target_size = (256, 256),
                            class_mode = None,
                            batch_size = 32,
                            seed=custom_seed
)

train_dataset_masks = data_generator.flow_from_directory(
                            directory=mask_directory,
                            target_size = (256, 256),
                            batch_size = 32,
                            class_mode = None,
                            color_mode = 'grayscale',
                            seed=custom_seed
)

train_generator = zip(train_dataset_images, train_dataset_masks)

运行时抛出ValueError: expected 1 input but received 2,尝试过两种合并生成器的函数均无效:

第一种合并函数:

def combine_generator(image_generator, mask_generator):
    while True:
        image_batch = image_generator.next()
        mask_batch = mask_generator.next()
        yield (image_batch, mask_batch)

第二种合并函数:

def combine_generator(image_gen, mask_gen):
    for img, mask in zip(image_gen, mask_gen):
        yield img, mask

提出以下疑问:

  1. 如何正确将ImageDataGenerator传入U-net?
  2. 图像与掩码是否必须同名?
  3. 掩码是否需按类别独热编码?

问题解答

1. 正确传入ImageDataGenerator到U-Net的方式

你遇到的错误本质是生成器返回的格式不符合U-Net的训练要求。segmentation_models的U-Net训练时,需要生成器返回**(输入图像批量, 目标掩码批量)**的元组结构,而你当前的生成器返回的是两个独立的批量数据,模型无法识别为“输入-标签”对。

修正后的合并生成器需要明确返回符合要求的结构,同时还要处理掩码的维度匹配模型输出:

import tensorflow as tf

def combine_generator(image_generator, mask_generator, num_classes):
    while True:
        img_batch = image_generator.next()
        mask_batch = mask_generator.next()
        # 若使用categorical_crossentropy损失,需转独热编码;若用sparse版本则跳过此步
        mask_batch = tf.keras.utils.to_categorical(mask_batch, num_classes=num_classes)
        # 返回(输入图像, 目标掩码)的标准训练格式
        yield img_batch, mask_batch

使用时传入你的类别数:

train_generator = combine_generator(train_dataset_images, train_dataset_masks, num_classes=你的类别数量)

另外要注意:两个flow_from_directory必须设置相同的seed,确保图像和掩码按相同顺序读取,保证配对正确;同时图像和掩码的目录结构要一致(比如都是一级目录下存放所有文件)。

2. 图像与掩码是否必须同名?

是的,必须同名(文件名主体一致,后缀可不同)。因为flow_from_directory是按文件名排序读取文件的,只有图像和掩码文件名完全匹配,相同seed下两个生成器才能取出对应的配对样本。如果文件名不一致,会导致图像和掩码错位,训练完全失效。

比如图像目录下有sample_001.jpg,掩码目录下必须有sample_001.png(或其他格式,只要文件名前缀一致),才能保证配对正确。

3. 掩码是否需按类别独热编码?

取决于你选用的损失函数和模型输出配置:

  • 若使用categorical_crossentropy损失函数,且模型最后一层用softmax激活输出num_classes通道的概率图,必须对掩码做独热编码,将单通道的类别索引(如0、1、2...)转为num_classes通道的二进制矩阵,每个通道对应一个类别的掩码。
  • 若使用sparse_categorical_crossentropy损失函数,模型输出同样是num_classes通道,但掩码可以保持单通道的类别索引格式,无需独热编码,这种方式更节省内存。

另外要注意:掩码中的像素值必须是连续的类别索引(比如0代表背景,1代表类别1,以此类推),不能是随机RGB值,否则模型无法正确学习类别边界。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 13:03:13