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

如何将ImageDataGenerator(DirectoryIterator)转换为tf.data.Dataset

报错原因

tf.data.Dataset.from_generator 要求传入的第一个参数是可调用对象(即函数/实现了__call__方法的对象),你直接传入了flow_from_directory返回的迭代器实例,不是可调用对象,因此触发generator must be a Python callable报错。
除此之外现有实现还有两个关键遗漏:

  • 没有将原始图像和分割掩码做配对,无法直接输入U-Net训练
  • 缺少掩码从RGB格式转二值分类标签的预处理步骤
修正实现

1. 拆分图像、掩码的数据增强配置

掩码不需要做基于数据集统计量的归一化,单独配置增强参数,仅保留几何变换类增强,保证和图像变换逻辑一致即可:

BATCH_SIZE = 8
INPUT_SHAPE = (256, 256, 3) # 替换为你的实际输入尺寸
seed = 123

# 图像生成器配置,保留归一化和几何增强
image_data_gen_args = dict(
    featurewise_center=True,
    featurewise_std_normalization=True,
    rotation_range=90,
    width_shift_range=0.1,
    height_shift_range=0.1,
    zoom_range=0.2
)
# 掩码生成器配置,仅保留几何增强,加固定rescale到0-1区间
mask_data_gen_args = dict(
    rotation_range=90,
    width_shift_range=0.1,
    height_shift_range=0.1,
    zoom_range=0.2,
    rescale=1/255.
)

image_datagen = ImageDataGenerator(**image_data_gen_args)
mask_datagen = ImageDataGenerator(**mask_data_gen_args)

# 初始化两个flow生成器,必须用相同seed保证增强对齐
image_generator = image_datagen.flow_from_directory(
    "image_directory/",
    classes=['images'],
    class_mode=None,
    target_size=INPUT_SHAPE[:2],
    batch_size=BATCH_SIZE,
    seed=seed
)
mask_generator = mask_datagen.flow_from_directory(
    "label_directory",
    classes=['labels'],
    class_mode=None,
    target_size=INPUT_SHAPE[:2],
    batch_size=BATCH_SIZE,
    seed=seed
)

# 提前拟合图像生成器的归一化统计量(均值、标准差)
sample_batch = next(image_generator)[0]
image_datagen.fit(np.expand_dims(sample_batch, 0), augment=True, seed=seed)

2. 编写可调用的配对生成函数

在函数内实现图像、掩码的迭代配对,同时嵌入二值掩码预处理逻辑:

import numpy as np

def binary_mask_preprocess(mask_rgb):
    # 输入是0-1区间的RGB掩码,输出单通道0/1二值标签
    # 如果你的前景色不是白色,修改这里的判断阈值即可
    binary_mask = np.where(mask_rgb.mean(axis=-1, keepdims=True) > 0.5, 1.0, 0.0)
    return binary_mask.astype(np.float32)

def train_generator():
    image_generator.reset()
    mask_generator.reset()
    for img_batch, mask_batch in zip(image_generator, mask_generator):
        processed_mask_batch = np.array([binary_mask_preprocess(m) for m in mask_batch])
        yield img_batch, processed_mask_batch

3. 转换为tf.data.Dataset对象

指定明确的输出签名,兼容TensorFlow最新版本接口:

ds = tf.data.Dataset.from_generator(
    train_generator,
    output_signature=(
        tf.TensorSpec(shape=(None, *INPUT_SHAPE), dtype=tf.float32),
        tf.TensorSpec(shape=(None, INPUT_SHAPE[0], INPUT_SHAPE[1], 1), dtype=tf.float32)
    )
)

# 加预取优化,提升训练吞吐量
ds = ds.prefetch(tf.data.AUTOTUNE)
注意事项
  • 两个flow_from_directory必须设置完全相同的随机种子,否则旋转、平移、缩放等几何增强无法在图像和掩码上对齐,训练会完全失效
  • 必须保证图像目录和掩码目录下的文件排序、文件名一一对应,否则会出现图像和掩码匹配错误
  • 由于Keras的flow_from_directory返回的是无限迭代的生成器,模型训练时需要手动指定steps_per_epoch = 12996 // BATCH_SIZE,避免单个epoch无限运行
  • 不要给掩码生成器开启featurewise_center、featurewise_std_normalization这类归一化操作,会破坏掩码的0/1分类取值

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 10:39:32