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

如何对TensorFlow加载的图像数据集执行交叉验证

基于tf.data.Dataset实现K折交叉验证的方法

你当前通过validation_split直接拆分得到的train_dataset和validation_dataset是单次留出法的拆分结果,无法直接基于这两个数据集执行交叉验证,需要先加载全量数据集再做K折拆分,具体操作如下:


步骤1:加载全量数据集并提取标签

首先不提前做训练/验证拆分,先拿到所有样本的基础信息,方便后续做K折划分,二分类任务推荐使用分层K折,可保证每折的正负样本比例和全集一致,避免分布偏移:

import tensorflow as tf
from sklearn.model_selection import StratifiedKFold
import numpy as np

# 先加载全量数据集,暂时不设分批、不拆分
dataset = tf.keras.preprocessing.image_dataset_from_directory(
    base_folder,
    label_mode='categorical',
    image_size=(img_height, img_width),
    batch_size=None,
    shuffle=True,
    seed=123
)

# 提取所有样本的标签,用于分层拆分
labels = []
for _, label in dataset:
    labels.append(np.argmax(label.numpy()))
labels = np.array(labels)

# 将数据集转为列表方便按索引取数
dataset_list = list(dataset.as_numpy_iterator())

步骤2:循环训练每折模型

# 定义折数,常用5折或10折
k = 5
skf = StratifiedKFold(n_splits=k, shuffle=True, random_state=123)

# 存储每折的验证结果
val_acc_list = []
val_loss_list = []

for fold, (train_idx, val_idx) in enumerate(skf.split(np.zeros(len(labels)), labels)):
    print(f"当前训练第 {fold+1}/{k} 折")
    
    # 按索引生成当前折的训练、验证数据集
    train_ds = tf.data.Dataset.from_generator(
        lambda: (dataset_list[i] for i in train_idx),
        output_types=(tf.float32, tf.float32),
        output_shapes=((img_height, img_width, 3), (2,)) # 灰度图通道数改为1即可
    )
    val_ds = tf.data.Dataset.from_generator(
        lambda: (dataset_list[i] for i in val_idx),
        output_types=(tf.float32, tf.float32),
        output_shapes=((img_height, img_width, 3), (2,))
    )
    
    # 分批、预取优化训练性能
    batch_size = 16
    train_ds = train_ds.shuffle(buffer_size=len(train_idx)).batch(batch_size).prefetch(tf.data.AUTOTUNE)
    val_ds = val_ds.batch(batch_size).prefetch(tf.data.AUTOTUNE)
    
    # 每次重新初始化模型,避免上一折权重干扰
    model = your_cnn_model() # 替换为你自己的CNN模型定义函数
    model.compile(
        optimizer='adam',
        loss='categorical_crossentropy',
        metrics=['accuracy']
    )
    
    # 训练模型
    history = model.fit(
        train_ds,
        validation_data=val_ds,
        epochs=10, # 替换为你自己的训练轮数
        verbose=1
    )
    
    # 记录当前折的验证结果
    val_loss, val_acc = model.evaluate(val_ds, verbose=0)
    val_acc_list.append(val_acc)
    val_loss_list.append(val_loss)
    print(f"第 {fold+1} 折验证准确率: {val_acc:.4f}, 验证损失: {val_loss:.4f}")

步骤3:统计交叉验证最终结果

print(f"\n{k}折交叉验证平均验证准确率: {np.mean(val_acc_list):.4f} ± {np.std(val_acc_list):.4f}")
print(f"{k}折交叉验证平均验证损失: {np.mean(val_loss_list):.4f} ± {np.std(val_loss_list):.4f}")

大内存优化提示:如果数据集体积很大,把全量数据转成列表会占用过多内存,可以直接读取数据集的文件路径做拆分,每折单独加载数据:

# 仅读取路径和标签,不加载图像到内存
full_ds = tf.keras.preprocessing.image_dataset_from_directory(base_folder, shuffle=False)
file_paths = full_ds.file_paths
labels = np.array(full_ds.class_names)[full_ds.labels]

后续K折拆分时直接对路径和标签做拆分,每折用tf.keras.utils.image_dataset_from_paths加载对应路径的数据即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 17:06:03