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

如何在Keras中协调图像与矩阵输入,实现样本对应匹配?

解决Keras多输入同步问题:图像与矩阵特征一一对应

这问题我之前处理过不少,核心思路是用自定义Sequence生成器把图像数据和矩阵特征绑定起来,从根源上保证每一组输入的对应关系,哪怕是批次训练也不会打乱顺序。下面给你一步步讲清楚怎么做:

1. 先确保原始数据的顺序对齐

在开始写代码前,你得先把图像路径列表和矩阵数据的行严格对应起来:

  • 比如你有n张图像,把它们的路径存在一个列表img_paths里,第i个元素就是第i张图的路径;
  • 你的矩阵数据是形状为(n, m)的NumPy数组matrix_data,第i行就是第i张图对应的m维特征。
  • 不管是从目录遍历图像,还是从CSV读取特征,一定要保证两者的索引完全匹配——这是后续同步的基础。

2. 自定义Multi-Input Sequence生成器

Keras的Sequence类专门用于构建有序的批次数据生成器,它会严格按照你定义的逻辑返回批次,不会像普通生成器那样出现顺序混乱的问题。我们继承这个类,实现一个能同时返回图像和矩阵特征的生成器:

import numpy as np
from tensorflow.keras.utils import Sequence
from tensorflow.keras.preprocessing.image import load_img, img_to_array

class MultiInputGenerator(Sequence):
    def __init__(self, img_paths, matrix_data, batch_size, img_size, shuffle=True, preprocess_fn=None):
        self.img_paths = img_paths
        self.matrix_data = matrix_data
        self.batch_size = batch_size
        self.img_size = img_size
        self.shuffle = shuffle
        self.preprocess_fn = preprocess_fn  # 可选:自定义图像预处理函数
        self.indexes = np.arange(len(self.img_paths))
        self.on_epoch_end()  # 初始化时打乱索引(如果需要)

    def __len__(self):
        # 计算每个epoch的总批次数
        return int(np.ceil(len(self.img_paths) / self.batch_size))

    def __getitem__(self, index):
        # 获取当前批次的索引范围
        batch_start = index * self.batch_size
        batch_end = min((index + 1) * self.batch_size, len(self.img_paths))
        batch_indexes = self.indexes[batch_start:batch_end]
        
        # 根据索引取对应的图像和矩阵特征
        batch_imgs = []
        for idx in batch_indexes:
            img = load_img(self.img_paths[idx], target_size=self.img_size)
            img = img_to_array(img)
            # 应用预处理(比如归一化、用预训练模型的预处理函数)
            if self.preprocess_fn:
                img = self.preprocess_fn(img)
            else:
                img = img / 255.0  # 默认归一化到0-1
            batch_imgs.append(img)
        
        batch_matrix = self.matrix_data[batch_indexes]
        # 返回多输入列表,如果有标签的话,改成 return [np.array(batch_imgs), batch_matrix], batch_labels
        return [np.array(batch_imgs), batch_matrix]

    def on_epoch_end(self):
        # 每个epoch结束后打乱索引,保证训练随机性,同时不破坏对应关系
        if self.shuffle:
            np.random.shuffle(self.indexes)

3. 初始化生成器并构建多输入模型

初始化生成器

假设你已经准备好了对齐的img_paths和matrix_data,直接实例化生成器即可:

# 示例数据:替换成你自己的实际路径和矩阵
img_paths = ["data/img_001.jpg", "data/img_002.jpg", ..., "data/img_n.jpg"]
matrix_data = np.random.rand(n, 10)  # n是样本数,10是特征数

train_generator = MultiInputGenerator(
    img_paths=img_paths,
    matrix_data=matrix_data,
    batch_size=32,
    img_size=(224, 224),
    shuffle=True
)

构建多分支融合模型

接下来构建一个包含图像CNN分支和矩阵全连接分支的模型,最后将两个分支的输出融合:

from tensorflow.keras.models import Model
from tensorflow.keras.layers import Input, Dense, Conv2D, MaxPooling2D, Flatten, concatenate

# 图像输入分支(用你自己的CNN结构即可)
img_input = Input(shape=(224, 224, 3))
x = Conv2D(32, (3,3), activation='relu')(img_input)
x = MaxPooling2D((2,2))(x)
x = Conv2D(64, (3,3), activation='relu')(x)
x = MaxPooling2D((2,2))(x)
x = Flatten()(x)
img_embedding = Dense(128, activation='relu')(x)

# 矩阵输入分支(根据你的特征数调整)
matrix_input = Input(shape=(10,))  # 这里的10对应matrix_data的特征数m
matrix_embedding = Dense(64, activation='relu')(matrix_input)
matrix_embedding = Dense(128, activation='relu')(matrix_embedding)

# 融合两个分支的特征
merged_features = concatenate([img_embedding, matrix_embedding])
# 输出层:根据你的任务调整(比如分类、回归)
output = Dense(1, activation='sigmoid')(merged_features)

# 定义多输入模型
model = Model(inputs=[img_input, matrix_input], outputs=output)
model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])

4. 开始训练

直接用model.fit()传入生成器即可,生成器会自动按批次返回同步的图像和矩阵特征:

model.fit(
    train_generator,
    epochs=15,
    # 如果有验证集,同样用MultiInputGenerator构建验证生成器传入
    # validation_data=val_generator
)

关键注意点

  • 顺序绝对不能乱:初始化img_paths和matrix_data时,一定要保证第i个路径对应第i行特征,哪怕是后续打乱,也是通过索引数组同时打乱两者的对应关系;
  • 预处理统一:如果图像用了预训练模型的预处理函数(比如tf.keras.applications.resnet50.preprocess_input),一定要在生成器里统一应用,避免数据不一致;
  • 验证集同理:验证集的生成器也要用同一个类构建,确保验证数据的对应关系也正确。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 08:33:23