如何用tf.keras.utils.Sequence创建多输入自定义数据生成器?
多输入自定义数据生成器实现(图像+结构化数据)
针对你的多输入模型(2个数值特征、1个分类特征、1个图像特征,输出二分类),我们可以通过继承Keras的Sequence类实现稳定的自定义数据生成器,解决大数据集无法一次性加载的问题。以下是具体实现方案:
核心逻辑
生成器的核心是分批加载并预处理不同类型的数据,每个批次需要同时返回三类结构化数据、图像数据,以及对应的二分类标签。其中图像数据需按需读取、预处理后送入CNN分支,结构化数据则直接提取或编码后送入全连接分支。
完整实现步骤
1. 依赖导入
import numpy as np import pandas as pd from tensorflow.keras.utils import Sequence from PIL import Image from sklearn.preprocessing import LabelEncoder from sklearn.model_selection import train_test_split from tensorflow.keras.models import Model from tensorflow.keras.layers import Input, Dense, Flatten, Conv2D, MaxPooling2D, concatenate
2. 自定义Sequence生成器类
class MultiInputGenerator(Sequence): def __init__(self, df, img_size, batch_size, cat_encoder, is_train=False): self.df = df.reset_index(drop=True) # 重置索引避免取数错误 self.img_height, self.img_width = img_size self.batch_size = batch_size self.cat_encoder = cat_encoder self.is_train = is_train # 标记是否为训练集,用于数据增强 self.labels = self.df['label'].values.astype('float32') def __len__(self): # 返回总批次数,向上取整确保所有数据都被覆盖 return int(np.ceil(len(self.df) / self.batch_size)) def __getitem__(self, index): # 生成单个批次的数据 start_idx = index * self.batch_size end_idx = min((index + 1) * self.batch_size, len(self.df)) batch_df = self.df.iloc[start_idx:end_idx] # 1. 处理数值特征 num_features = batch_df[['num_feature1', 'num_feature2']].values.astype('float32') # 2. 处理分类特征(使用提前拟合好的编码器) cat_features = self.cat_encoder.transform(batch_df['cat_feature']).reshape(-1, 1) # 3. 处理图像特征:批量读取、预处理 img_features = [] for img_path in batch_df['image_path']: # 读取图像并调整尺寸 img = Image.open(img_path).resize((self.img_width, self.img_height)) # 归一化到0-1区间 img = np.array(img) / 255.0 # 训练集可选数据增强(示例:随机水平翻转) if self.is_train and np.random.rand() > 0.5: img = np.fliplr(img) img_features.append(img) img_features = np.array(img_features) # 4. 获取当前批次标签 batch_labels = self.labels[start_idx:end_idx] # 返回多输入列表和标签 return [num_features, cat_features, img_features], batch_labels def on_epoch_end(self): # 每个epoch结束后打乱训练集数据,避免顺序影响 if self.is_train: self.df = self.df.sample(frac=1).reset_index(drop=True)
3. 预处理分类特征
提前对分类特征进行编码,确保训练和验证集使用同一套编码规则:
# 假设你的数据框名为df,分类特征列名为cat_feature cat_encoder = LabelEncoder() cat_encoder.fit(df['cat_feature']) # 如果需要独热编码,可替换为OneHotEncoder: # from sklearn.preprocessing import OneHotEncoder # cat_encoder = OneHotEncoder(sparse_output=False) # cat_encoder.fit(df[['cat_feature']]) # 注意:此时分类特征的输入shape要对应独热编码后的维度
4. 构建多输入模型
# 定义各输入分支 # 数值输入分支 num_input = Input(shape=(2,), name='numeric_input') num_x = Dense(32, activation='relu')(num_input) num_x = Dense(16, activation='relu')(num_x) # 分类输入分支(LabelEncoder编码后shape为(1,)) cat_input = Input(shape=(1,), name='categorical_input') cat_x = Dense(16, activation='relu')(cat_input) # 图像输入分支(假设输入为224x224的RGB图像) img_input = Input(shape=(224, 224, 3), name='image_input') img_x = Conv2D(32, (3, 3), activation='relu')(img_input) img_x = MaxPooling2D((2, 2))(img_x) img_x = Conv2D(64, (3, 3), activation='relu')(img_x) img_x = MaxPooling2D((2, 2))(img_x) img_x = Conv2D(128, (3, 3), activation='relu')(img_x) img_x = MaxPooling2D((2, 2))(img_x) img_x = Flatten()(img_x) img_x = Dense(128, activation='relu')(img_x) # 拼接所有分支特征 merged = concatenate([num_x, cat_x, img_x]) merged = Dense(64, activation='relu')(merged) # 二分类输出层 output = Dense(1, activation='sigmoid', name='binary_output')(merged) # 构建模型 model = Model(inputs=[num_input, cat_input, img_input], outputs=output) model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])
5. 使用生成器训练模型
# 拆分训练集和验证集 train_df, val_df = train_test_split(df, test_size=0.2, random_state=42) # 创建训练/验证生成器 train_gen = MultiInputGenerator( df=train_df, img_size=(224, 224), batch_size=32, cat_encoder=cat_encoder, is_train=True ) val_gen = MultiInputGenerator( df=val_df, img_size=(224, 224), batch_size=32, cat_encoder=cat_encoder, is_train=False ) # 启动训练,开启多进程加速加载 model.fit( train_gen, validation_data=val_gen, epochs=15, workers=4, use_multiprocessing=True )
关键注意事项
- 图像预处理一致性:训练和验证集的图像必须使用相同的resize、归一化规则,数据增强仅应用于训练集。
- 编码器复用:分类特征的编码器必须用训练集拟合,再用于验证集,避免数据泄露。
- 多进程支持:使用
Sequence类而非普通生成器,可安全开启多进程加速图像读取,提升训练效率。
内容的提问来源于stack exchange,提问作者Abdullah Al Munem
相关产品推荐
相关产品推荐

