基于ResNet-50的多输入模型:如何融合图像与附加文本数据?
解决多输入(图像+附加特征)的ResNet-50模型训练问题
首先,你面临的核心问题是Keras的flow_from_directory只能生成图像数据,没法同时加载CSV里的附加特征,所以自定义Sequence生成器绝对是最靠谱的解决方案——你已经找对方向了!下面我会帮你梳理这个方案的关键细节,还能优化一下你现有的实现:
为什么选Sequence而不是普通生成器?
Sequence是Keras官方推荐的批量数据加载工具,比普通生成器好用太多:
- 支持多进程训练,稳定性拉满(不会出现数据重复/遗漏的问题)
- 自带索引管理,训练流程和Keras的契合度更高
- 代码结构更清晰,后期维护起来更省心
你的自定义Sequence优化版
你的实现已经覆盖了核心逻辑,我补充几个关键优化点和细节处理:
1. 拆分预处理逻辑,让代码更清爽
把图像和CSV数据的预处理拆成独立函数,方便后期调整:
def preprocess_image(image_path, target_size, color_mode="grayscale"): img = Image.open(image_path).convert(color_mode) img = img.resize((target_size, target_size)) img_array = np.array(img) # 和你之前用的rescale保持一致,做归一化 img_array = img_array / 255.0 # 单通道图像要补上通道维度 if color_mode == "grayscale": img_array = np.expand_dims(img_array, axis=-1) return img_array def preprocess_csv_row(csv_df, filename, label_dict, num_classes): # 从CSV中定位当前图像的行数据 row = csv_df.loc[filename] # 处理多标签:按分隔符拆分后转成one-hot编码 labels = row["Labels"].split("|") one_hot = np.zeros(num_classes) for label in labels: if label in label_dict: one_hot[label_dict[label]] = 1.0 # 提取附加特征(这里对应CSV里的Info1和Info2,vec_size=2) features = np.array([row["Info1"], row["Info2"]]) return features, one_hot
2. 优化Sequence类的核心逻辑
避免重复读取CSV,提前做好索引映射,同时处理路径安全拼接:
import os import numpy as np import pandas as pd from PIL import Image from tensorflow.keras.utils import Sequence class CustomSequenceGenerator(Sequence): def __init__(self, image_dir, csv_file_path, label_list, dim=448, batch_size=8, n_classes=15, n_channels=1, vec_size=2, shuffle=True): self.image_dir = image_dir # 只保留CSV里存在的图像文件,避免出现找不到对应数据的情况 self.csv_df = pd.read_csv(csv_file_path).set_index("File", drop=True) self.image_file_list = [f for f in os.listdir(image_dir) if f in self.csv_df.index] self.batch_size = batch_size self.n_classes = n_classes self.dim = dim self.n_channels = n_channels self.shuffle = shuffle self.vec_size = vec_size self.label_dict = dict(zip(label_list, range(len(label_list)))) # 初始化索引数组,用于打乱数据 self.indexes = np.arange(len(self.image_file_list)) if self.shuffle: np.random.shuffle(self.indexes) def __len__(self): # 计算总批次数,向上取整保证所有数据都能用到 return int(np.ceil(len(self.image_file_list) / self.batch_size)) def __getitem__(self, index): # 获取当前批次的索引 batch_indexes = self.indexes[index*self.batch_size : (index+1)*self.batch_size] # 根据索引拿到对应的图像文件名 batch_samples = [self.image_file_list[i] for i in batch_indexes] # 生成当前批次的训练数据 return self.__data_generation(batch_samples) def on_epoch_end(self): # 每个epoch结束后打乱索引,保证训练数据的随机性 if self.shuffle: np.random.shuffle(self.indexes) def __data_generation(self, batch_samples): # 初始化批次数据数组(注意用len(batch_samples)而不是固定batch_size,避免最后一批数据不足的问题) x_image = np.empty((len(batch_samples), self.dim, self.dim, self.n_channels)) x_vec = np.empty((len(batch_samples), self.vec_size)) y = np.empty((len(batch_samples), self.n_classes)) for i, filename in enumerate(batch_samples): # 加载并预处理图像 img_path = os.path.join(self.image_dir, filename) x_image[i] = preprocess_image(img_path, self.dim, color_mode="grayscale" if self.n_channels ==1 else "rgb") # 加载并预处理CSV里的附加特征和标签 x_vec[i], y[i] = preprocess_csv_row(self.csv_df, filename, self.label_dict, self.n_classes) return [x_image, x_vec], y
3. 如何用这个生成器训练模型?
假设你的双输入模型已经定义好,直接把生成器传入model.fit()就行:
# 示例:从文件读取标签列表 def get_class_labels(label_path): with open(label_path, "r") as f: return [line.strip() for line in f.readlines()] train_labels = get_class_labels("path/to/your/labels.txt") # 创建训练和验证生成器 train_generator = CustomSequenceGenerator( image_dir="root/train", csv_file_path="path/to/train_data.csv", label_list=train_labels, dim=448, batch_size=8, n_classes=15, n_channels=1, vec_size=2, shuffle=True ) val_generator = CustomSequenceGenerator( image_dir="root/validation", csv_file_path="path/to/val_data.csv", label_list=train_labels, dim=448, batch_size=8, n_classes=15, n_channels=1, vec_size=2, shuffle=False # 验证集不需要打乱数据 ) # 训练模型 history = model.fit( train_generator, epochs=10, validation_data=val_generator, use_multiprocessing=True, # 可选,根据CPU核心数开启,加快训练 workers=4 )
重要提示:多标签分类的适配
看到你CSV里有class1|class2这种多标签数据,这里要提醒一句:
如果你的任务是多标签分类,模型输出层的激活函数应该用
sigmoid而不是softmax,损失函数改用binary_crossentropy。因为softmax会强制所有类别概率和为1,完全不适合多标签场景哦!
内容的提问来源于stack exchange,提问作者Predrag Antanasijevic
相关产品推荐
相关产品推荐

