如何对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
相关产品推荐
相关产品推荐

