如何用TensorFlow内置方法实现Sklearn风格的特征与标签张量划分
使用TensorFlow实现类似sklearn的训练测试集划分
嘿,我懂你想要的是什么——直接用TensorFlow原生工具处理tf.Tensor格式的特征和标签,实现和sklearn.model_selection.train_test_split完全等价的无交集划分,对吧?其实TensorFlow本身没有同名的API,但我们用它内置的方法就能轻松搞定,下面给你两种实用方案:
方案一:用tf.data.Dataset.split(TensorFlow 2.13+推荐)
从TF 2.13版本开始,tf.data.Dataset新增了split方法,按比例划分数据集的操作变得非常直观,而且完全基于TensorFlow原生实现:
import tensorflow as tf # 第一步:把特征和标签打包成tf.data.Dataset对象 dataset = tf.data.Dataset.from_tensor_slices((features, labels)) # 按比例拆分,这里测试集占10%,设置seed保证结果可复现 train_dataset, test_dataset = dataset.split([0.9, 0.1], seed=123) # 如果需要把Dataset转回tf.Tensor格式(和你想要的返回形式一致): X_train, y_train = next(iter(train_dataset.batch(len(train_dataset)))) X_test, y_test = next(iter(test_dataset.batch(len(test_dataset))))
方案二:手动索引划分(兼容所有TF 2.x版本)
如果你还在使用较早的TF版本,手动生成随机索引的方法兼容性拉满,同样能保证数据集无重叠:
import tensorflow as tf # 获取总样本数量 total_samples = features.shape[0] # 生成随机打乱的索引序列,设置seed确保划分结果固定 shuffled_indices = tf.random.shuffle(tf.range(total_samples), seed=123) # 计算训练集和测试集的样本量 test_ratio = 0.1 test_sample_count = int(total_samples * test_ratio) train_sample_count = total_samples - test_sample_count # 拆分索引 train_indices = shuffled_indices[:train_sample_count] test_indices = shuffled_indices[train_sample_count:] # 根据索引提取对应的特征和标签 X_train = tf.gather(features, train_indices) y_train = tf.gather(labels, train_indices) X_test = tf.gather(features, test_indices) y_test = tf.gather(labels, test_indices)
这两种方法都能实现你要的效果:随机划分、训练测试集无交集、支持设置随机种子复现结果。第一种写法更简洁,适合新版本;第二种则能适配所有TensorFlow 2.x环境,按需选择就好。
内容的提问来源于stack exchange,提问作者Hawklaz
相关产品推荐
相关产品推荐

