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

如何在UNet分类任务中用flow_from_directory处理多路径多类别数据

问题与解决方案

问题点

  • 要拼接6张GLCM图像,用7类共42个路径的数据做UNet分类,但训练时只读取class1的掩码,还卡在for six_img, six_mask in zip(img_zip,mask_zip):循环出不来
  • 没法让image_generator和mask_generator同时读取多路径,传路径列表不被支持

原代码核心问题

  1. 遍历folders时每次覆盖生成器变量,根本没收集每6个一组的生成器
  2. zip(*img_append)用法错误:生成器是迭代器,直接拆包zip会导致逻辑混乱,无法正确获取同批次的6张图
  3. 先把所有拼接后的图像存到列表再yield,既占内存又会导致生成器提前耗尽,无法持续生成训练数据

修正后的代码

import numpy as np
from tensorflow.keras.preprocessing.image import ImageDataGenerator

folders = [
    'classes1/glcm1',
    'classes1/glcm2',
    'classes1/glcm3',
    'classes1/glcm4',
    'classes1/glcm5', 
    'classes1/glcm6',
    'classes2/glcm1',
    # ... 省略剩余42个路径
    'classes7/glcm6',
]

def glcm_unet_generator(folders, aug_dict, batch_size=1, target_size=(512,512)):
    # 按每6个路径为一组(对应同一类的6个GLCM图像)分组
    grouped_folders = [folders[i:i+6] for i in range(0, len(folders), 6)]
    
    for group in grouped_folders:
        image_generators = []
        mask_generators = []
        # 给每组生成器设置相同seed,保证数据增强同步
        seed = np.random.randint(1000)
        for path in group:
            img_gen = ImageDataGenerator(**aug_dict).flow_from_directory(
                path,
                target_size=target_size,
                batch_size=batch_size,
                class_mode=None,
                seed=seed,
                shuffle=True
            )
            mask_gen = ImageDataGenerator(**aug_dict).flow_from_directory(
                path,  # 若掩码和图像路径不同,需单独修改此处路径
                target_size=target_size,
                batch_size=batch_size,
                class_mode=None,
                seed=seed,
                shuffle=True,
                color_mode='grayscale'  # 掩码一般为单通道,根据实际情况调整
            )
            image_generators.append(img_gen)
            mask_generators.append(mask_gen)
        
        # 持续生成该组的拼接数据,不会终止循环
        while True:
            # 获取每个生成器的当前批次数据
            batch_imgs = [next(gen) for gen in image_generators]
            batch_masks = [next(gen) for gen in mask_generators]
            
            # 按通道维度拼接6张GLCM图像
            concatenated_img = np.concatenate(batch_imgs, axis=3)
            # 掩码同理拼接
            concatenated_mask = np.concatenate(batch_masks, axis=3)
            
            yield (concatenated_img, concatenated_mask)

关键说明

  • 先把42个路径按每6个一组划分,对应同一类的6个GLCM图像
  • 同组生成器用相同seed,确保图像和掩码的增强操作同步(比如同时翻转、平移)
  • 用while True持续生成批次,避免循环提前终止
  • 实时拼接同批次数据,不占用额外内存,适配大规模训练场景
  • 若图像和掩码存储路径不同,需分别调整图像、掩码生成器的path参数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 06:02:03