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

TensorFlow/Keras双输入模型的Generator训练及验证适配问题

解决Keras Sequence生成器用于ArcFace模型验证的问题

核心结论

验证数据生成器的逻辑和训练生成器完全一致,只需保证__getitem__的返回格式与训练时相同:( (X_val, y_val), y_val ),就能直接传入model.fit的validation_data参数。

具体实现步骤

  1. 复用/调整训练生成器类
    验证生成器可直接复用训练用的DataGenerator类,初始化时传入验证数据集的索引、路径等信息,同时关闭数据增强(如果训练时使用了),并将shuffle设为False(避免打乱验证数据顺序,方便后续对应评估)。

示例代码:

import numpy as np
from tensorflow import keras

class DataGenerator(keras.utils.Sequence):
    def __init__(self, data_indices, labels, batch_size=32, shuffle=True, augment=False):
        self.data_indices = data_indices
        self.labels = labels
        self.batch_size = batch_size
        self.shuffle = shuffle
        self.augment = augment
        self.on_epoch_end()

    def __len__(self):
        return int(np.ceil(len(self.data_indices) / self.batch_size))

    def __getitem__(self, index):
        # 取出当前batch的索引范围
        batch_indices = self.data_indices[index*self.batch_size : (index+1)*self.batch_size]
        # 加载对应batch的特征X和标签y
        X = self.load_features(batch_indices)
        y = self.labels[batch_indices]

        # 验证集跳过数据增强
        if self.augment:
            X = self.apply_augmentation(X)

        # 返回符合模型要求的格式:(输入张量组, 损失计算用标签)
        return (X, y), y

    def on_epoch_end(self):
        if self.shuffle:
            np.random.shuffle(self.data_indices)

    # 自定义特征加载逻辑,根据你的数据存储方式实现
    def load_features(self, indices):
        # 从磁盘/数据库加载对应索引的特征数据
        pass

    # 自定义训练时的数据增强逻辑
    def apply_augmentation(self, X):
        # 对特征X做增强处理
        pass
  1. 实例化验证生成器
# 假设val_indices是验证数据的索引列表,val_labels是对应标签数组
val_generator = DataGenerator(
    data_indices=val_indices,
    labels=val_labels,
    batch_size=32,
    shuffle=False,
    augment=False
)
  1. 接入训练流程
    在model.fit中直接指定validation_data=val_generator即可:
model.fit(
    train_generator,
    epochs=10,
    validation_data=val_generator,
    # 按需设置steps_per_epoch、validation_steps等参数
)

格式说明

  • (X, y)是模型要求的两个输入张量,对应你定义的模型输入[X, y],满足模型对输入数量的要求。
  • 第二个y是ArcFace损失计算、验证指标评估所需的真实标签,和训练时的逻辑完全匹配。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 16:15:26