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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 12:23:12