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

torch.utils.data.random_split与numpy随机划分的性能差异原因问询

问题:PyTorch random_split与手动numpy划分K折CV的性能差异

我用PyTorch Lightning实现CNN,需要用5折交叉验证评估跨数据集性能(已注意数据泄露和过拟合问题)。但发现两种划分方法的测试AUC差异显著:

  • 常规8:1:1训练/验证/测试划分下,测试集AUC约0.95
  • 手动numpy实现的5折CV,AUC仅0.88
  • PyTorch内置random_split实现的5折CV,AUC能达到0.95

测试样本量约10000,两种方法的AUC差异具有统计显著性。想搞清楚:
为什么random_split的结果更好?我的numpy实现里样本属于某一折的概率是独立的,random_split是同样的逻辑还是采用了随机区间划分的方式?


K折实验代码

手动numpy划分(AUC=0.88)

fold    = 1
indices = np.arange(X.shape[0])
sampler = np.random.permutation(indices) % 5 # 5 fold CV
X_train, X_test = X[(sampler!=fold)], X[(sampler==fold)]
y_train, y_test = y[(sampler!=fold)], y[(sampler==fold)]

dataset = TensorDataset(torch.Tensor(X_train), torch.Tensor(y_train))
train_dataset, val_dataset = torch.utils.data.random_split(dataset, [0.9, 0.1])
test_dataset  = TensorDataset(torch.Tensor(X_test), torch.Tensor(y_test))

PyTorch random_split划分(AUC=0.95)

fold = 1
dataset       = TensorDataset(torch.Tensor(X), torch.Tensor(y))
five_folds    = torch.utils.data.random_split(dataset, [1/5]*5) # 5 fold CV
test_dataset  = five_folds[fold]
train_dataset = torch.utils.data.ConcatDataset(
    [five_folds[i] for i in range(5) if i != fold]
)
train_dataset, val_dataset = torch.utils.data.random_split(train_dataset, [0.9, 0.1])

原因分析

首先明确random_split的核心逻辑:它做的是无放回的随机划分——先对整个数据集的索引做一次全局随机排列,再按照你输入的比例/长度切割成连续区间。每个样本只会被分到一个折里,且各折的样本量严格符合预期(比如5折的话每折样本量接近总样本数的1/5)。

你的numpy实现存在两个关键问题:

  1. 样本量不均:np.random.permutation(indices) %5是给每个随机排列后的样本独立分配0-4的标签,这种方式下各折的样本量可能偏差极大(尤其是总样本数不是5的整数倍时)。比如总样本数N如果不是5的倍数,有的折可能远少于1/5的样本,有的远多于,这会直接导致训练集和测试集的数据分布偏离原数据集,模型性能自然下降。
  2. 类别分布失衡:如果你的数据集存在类别不平衡,独立分配标签的方式可能让某一折的类别占比严重偏离整体(比如某折里某类样本占比过高),而random_split的无放回随机划分能更好地保留原数据集的类别分布,保证测试集的代表性。

另外,检查你的numpy代码:y[(sampler!=fold]少了一个右括号,要是实际运行时没修正,会导致数据加载错误,这也可能是性能差的直接原因。

总结:random_split是无放回的随机划分,能保证各折样本量和分布接近原数据集;而你的numpy实现是独立分配标签,本质是有放回式的随机分配,会引发样本量和分布失衡,最终拉低模型性能。


内容的提问来源于stack exchange,提问作者JanEgern

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 18:33:30