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

使用tf.keras ImageDataGenerator进行语义分割:保持标签图像的整数类别值

解决ImageDataGenerator处理语义分割标签时数值被修改的问题

你的问题核心在于ImageDataGenerator的增强操作默认使用双线性插值,即使你在flow_from_directory里设置了interpolation="nearest",这个参数只作用于图像resize环节,而非旋转、平移等增强步骤。同时,生成器内部会自动将数据转为float32,覆盖了你设置的dtype="uint8",最终导致标签的整数类别被破坏。

下面给出几种基于你现有代码的简洁解决方案,帮你继续使用ImageDataGenerator同时保留标签的完整性:

方案1:批量后处理修正标签(最简单)

直接在zip后的生成器里对每个批次的标签做取整和类型转换,修正增强带来的浮点误差:

def corrected_train_generator(X_gen, y_gen):
    for X_batch, y_batch in zip(X_gen, y_gen):
        # 取整并转回uint8,恢复原始整数类别
        y_batch = np.round(y_batch).astype(np.uint8)
        yield X_batch, y_batch

# 替换原来的train_generator
train_generator = corrected_train_generator(X_gen, y_gen)

验证一下:

sample_image, sample_mask = next(train_generator)
print(f"掩码数据类型: {sample_mask.dtype}")
print(f"掩码唯一值: {np.unique(sample_mask)}")

你会看到掩码回到uint8类型,且所有值都是整数类别,没有小数。

方案2:自定义生成器类,强制增强时用最近邻插值

如果你想从根源避免浮点插值,可以继承ImageDataGenerator,重写增强变换的插值方式:

class MaskDataGenerator(ks.preprocessing.image.ImageDataGenerator):
    def apply_transform(self, x, transform_parameters):
        # 对掩码的所有空间变换(旋转、平移等)强制使用最近邻插值
        return super().apply_transform(x, transform_parameters, interpolation='nearest')

然后用这个类创建标签生成器:

y_generator = MaskDataGenerator(
    **args_aug, 
    cval=NoDataValue,
    dtype=np.uint8
)
y_gen = y_generator.flow_from_directory(
    directory="/path/to/y", 
    color_mode="grayscale", 
    interpolation="nearest", 
    **args_flow
)

train_generator = zip(X_gen, y_gen)

这种方式让增强操作本身就不会产生浮点值,从源头解决问题。

额外注意事项

  • 确保X和y生成器的seed完全一致(你已经设置了seed=42),这是保证图像和掩码增强同步的关键。
  • 标签生成器的fill_mode="constant"和cval=NoDataValue设置是正确的,避免填充时引入无效类别值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 17:52:40