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

如何拆分数据集用于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 17:42:02