如何拆分数据集用于CNN验证?tf.data.Dataset调用train_test_split报错怎么解决
问题解决方法
报错原因
sklearn.model_selection.train_test_split仅支持Numpy数组、Python列表等常规可迭代集合类型,不支持直接处理tf.data.Dataset类对象,因此会触发类型错误。
解决方案
方案1:先拆分Numpy数组再构建数据集(更推荐)
该方案可以支持分层抽样,保证拆分后训练集和验证集的正负样本比例和原数据集一致,避免类别分布偏移:
import numpy as np from sklearn.model_selection import train_test_split import tensorflow as tf # 合并所有样本和标签 X = np.concatenate([arraynegativos, arraypositivos], axis=0) y = np.concatenate([arrayceros, arrayunos], axis=0) # 按8:2比例拆分,stratify参数保证分层抽样 X_train, X_val, y_train, y_val = train_test_split(X, y, test_size=0.2, random_state=25, stratify=y) # 分别构建训练和验证数据集 trainingdataset = tf.data.Dataset.from_tensor_slices((X_train, y_train)).shuffle(len(X_train)) validatedataset = tf.data.Dataset.from_tensor_slices((X_val, y_val))
方案2:直接用tf.data API拆分数据集
如果不想转换为数组,可以直接用TensorFlow内置的数据集拆分方法,适合数据量过大无法全部加载到内存的场景:
# 计算训练集、验证集样本量 train_size = int(0.8 * n_total_img) val_size = n_total_img - train_size # 从已打乱的全量数据集中拆分 trainingdataset = datasetfinal.take(train_size) validatedataset = datasetfinal.skip(train_size)
注意:使用该方案需要保证原数据集已经完成充分打乱,你代码中已经调用
shuffle(n_total_img)符合要求。如果你的数据集类别不平衡,优先选择方案1。
内容的提问来源于stack exchange,提问作者Javier Decena Castillo
相关产品推荐
相关产品推荐

