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

如何在Keras模型中适配自定义简易图像生成器?

如何将自定义图像生成器适配到Keras模型

嘿,我来帮你搞定这个适配问题!首先得明确:Keras的model.fit()对数据生成器有几个硬性要求,咱们对照你的现有代码一步步调整就行。

你的my_iterator目前是单样本循环,而且是无限迭代,没法直接给Keras用——Keras需要生成器每次返回一个批量的输入+对应标签,还要能配合训练周期(epoch)控制迭代逻辑。下面给你两种实用方案,按需选就行:


方案1:改造现有生成器为Keras兼容版

直接在你现有代码基础上,加上批量处理、epoch打乱和数据格式化逻辑:

from PIL import Image, ImageOps
import numpy as np

def my_data_generator(batch_size, train_df, img_dir='master_train/'):
    num_samples = len(train_df)
    while True:
        # 每个epoch打乱一次数据,保证训练随机性
        train_df = train_df.sample(frac=1).reset_index(drop=True)
        # 按批量截取数据
        for i in range(0, num_samples, batch_size):
            batch_df = train_df[i:i+batch_size]
            batch_imgs = []
            batch_labels = []
            
            for _, row in batch_df.iterrows():
                # 读取图像(补全你没写完的处理逻辑)
                img = Image.open(f'{img_dir}{row["Image"]}').convert('L')
                longer_side = max(img.size)
                # 计算填充并补边(转整数,PIL要求参数是整数)
                horizontal_padding = int((longer_side - img.size[0]) / 2)
                vertical_padding = int((longer_side - img.size[1]) / 2)
                img = ImageOps.pad(img, (longer_side, longer_side), color=0)
                
                # 转成numpy数组,做归一化(根据你的模型需求调整,比如除以255)
                img_array = np.array(img) / 255.0
                # 给灰度图加通道维度(Keras输入一般是4D:样本数, H, W, 通道数)
                img_array = np.expand_dims(img_array, axis=-1)
                
                batch_imgs.append(img_array)
                batch_labels.append(row['Id'])
            
            # 转成Keras能识别的numpy数组格式
            yield (np.array(batch_imgs), np.array(batch_labels))

训练时这么用:

batch_size = 32
num_epochs = 10
# 每个epoch需要跑的步数:总样本数//批量大小
steps_per_epoch = len(train_df) // batch_size

# 初始化生成器
train_generator = my_data_generator(batch_size, train_df)

# 喂给模型训练
model.fit(
    train_generator,
    epochs=num_epochs,
    steps_per_epoch=steps_per_epoch
)

方案2:用keras.utils.Sequence(更推荐!)

Keras专门提供了Sequence类来处理批量数据生成,自带线程安全,还能自动处理epoch的遍历逻辑,比自定义生成器省心太多:

from PIL import Image, ImageOps
import numpy as np
from tensorflow.keras.utils import Sequence

class MyImageSequence(Sequence):
    def __init__(self, df, batch_size, img_dir='master_train/', fixed_img_size=None, normalize=True):
        self.df = df.copy()
        self.batch_size = batch_size
        self.img_dir = img_dir
        self.fixed_img_size = fixed_img_size  # 可以指定固定尺寸,不用动态计算长边
        self.normalize = normalize
        # 如果是分类任务,提前做标签映射(示例)
        # self.label_map = {label: idx for idx, label in enumerate(df['Id'].unique())}
        # self.num_classes = len(self.label_map)

    def __len__(self):
        # 返回每个epoch的训练步数,Keras会自动用这个值控制迭代
        return len(self.df) // self.batch_size

    def __getitem__(self, idx):
        # 返回第idx个批量的数据
        batch_start = idx * self.batch_size
        batch_end = batch_start + self.batch_size
        batch_df = self.df.iloc[batch_start:batch_end]
        
        batch_imgs = []
        batch_labels = []
        
        for _, row in batch_df.iterrows():
            img_path = f'{self.img_dir}{row["Image"]}'
            img = Image.open(img_path).convert('L')
            
            # 处理图像尺寸
            if self.fixed_img_size:
                img = img.resize((self.fixed_img_size, self.fixed_img_size))
            else:
                longer_side = max(img.size)
                img = ImageOps.pad(img, (longer_side, longer_side), color=0)
            
            # 预处理转数组
            img_array = np.array(img)
            if self.normalize:
                img_array = img_array / 255.0
            # 加通道维度
            img_array = np.expand_dims(img_array, axis=-1)
            
            batch_imgs.append(img_array)
            # 处理标签:如果是分类任务,转成one-hot编码(示例)
            # label_idx = self.label_map[row['Id']]
            # batch_labels.append(tf.keras.utils.to_categorical(label_idx, self.num_classes))
            batch_labels.append(row['Id'])
        
        return np.array(batch_imgs), np.array(batch_labels)

    def on_epoch_end(self):
        # 每个epoch结束时自动打乱数据,增强训练随机性
        self.df = self.df.sample(frac=1).reset_index(drop=True)

使用起来更简洁:

batch_size = 32
train_sequence = MyImageSequence(train_df, batch_size)

# 直接喂给模型,不用指定steps_per_epoch
model.fit(
    train_sequence,
    epochs=10
)

几个关键注意事项

  • 输入形状匹配:确保生成的图像形状和模型输入层完全一致!比如模型输入是(None, 224, 224, 1),那每个图像就得是(224,224,1)(灰度图的通道数是1)。
  • 标签格式适配:如果是分类任务,要根据模型输出层调整标签:比如输出是softmax多分类,标签要转成one-hot;输出是sigmoid二分类,标签是0/1数值。
  • 数据归一化:一定要把像素值缩放到合适区间(比如[0,1]),和模型训练的预处理逻辑保持一致。
  • 线程安全:Sequence是线程安全的,适合开启多线程训练(fit()里加workers参数);自定义生成器如果要多线程,得注意数据打乱的随机性问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 09:18:55