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

如何将sklearn train_test_split输出转为PyTorch DataLoader用于二分类任务

结论

完全不需要弃用现有的train_test_split拆分逻辑,你已经完成了数据加载、拆分的核心步骤,只需要额外写一个极简的自定义Dataset对接拆分结果,就能快速生成所需的DataLoader。

实现步骤

  1. 导入PyTorch相关依赖
    首先补充导入Dataset、DataLoader和torch工具库:
import torch
from torch.utils.data import Dataset, DataLoader
  1. 自定义适配内存数据的Dataset类
    因为你已经把所有图像、标签加载到内存的列表中,不需要在Dataset中重复读取本地文件,直接写一个接收X和y列表的通用Dataset即可:
class CustomImgDataset(Dataset):
    def __init__(self, x_list, y_list):
        # 接收拆分好的图像列表、标签列表
        self.x = x_list
        self.y = y_list
    
    def __len__(self):
        # 返回数据集总样本数
        return len(self.x)
    
    def __getitem__(self, idx):
        # 取对应下标的样本,转为PyTorch要求的格式
        img = self.x[idx]
        # 调整维度:从OpenCV输出的(H,W,C)转为PyTorch要求的(C,H,W)
        # 转为float32张量,同时归一化到0-1区间(可选但推荐)
        img_tensor = torch.tensor(img).permute(2,0,1).float() / 255.0
        # 标签转为长整型张量适配分类任务损失函数
        label_tensor = torch.tensor(self.y[idx]).long()
        return img_tensor, label_tensor
  1. 实例化Dataset并生成DataLoader
    在你原有拆分代码的基础上,追加以下代码即可得到训练、测试的DataLoader:
# 实例化训练集、测试集Dataset
train_dataset = CustomImgDataset(x_train, y_train)
test_dataset = CustomImgDataset(x_test, y_test)

# 生成DataLoader,可根据需求调整batch_size、shuffle、num_workers等参数
train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, num_workers=0)
test_loader = DataLoader(test_dataset, batch_size=16, shuffle=False, num_workers=0)

原有代码优化建议

你原有代码中提取标签的逻辑依赖Windows路径分隔符\,跨平台运行时会出错,建议替换为通用的文件名提取方式:

import os
# 替换原有标签提取逻辑
y.append(int(os.path.basename(name).split('_')[1]))

补充说明

如果你的数据集体量非常大,全部加载到内存会导致显存/内存溢出,才需要考虑调整逻辑:把文件路径按拆分结果传给Dataset,在__getitem__中实时读取本地图像。你当前的逻辑对于中小体量数据集完全可用,无需修改。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 02:09:02