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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 08:10:56