ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

深度学习数据管线实战:datasets与torch.utils.data高效协同

深度学习数据管线实战:datasets与torch.utils.data高效协同 我一直觉得深度学习项目里最容易被低估、又最能在关键时刻拖垮进度的环节就是数据处理。模型结构可以照着论文抄训练参数可以照着别人的配置调唯独数据处理每个数据集都有自己的脾气每个框架的加载方式也各有各的毛病。早几年我在跑图像分类和文本分类的时候被DataLoader和数据集格式转换折腾过很多次后来慢慢把Hugging Face的datasets和PyTorch自带的torch.utils.data配合起来用才算把这条链路彻底理顺了。这篇文章就把我在这套组合上的实战经验写透。内容包括两个数据栈各自的定位、分界线在哪里datasets库怎么把原始文件变成干净的数据集torch.utils.data底层到底怎么工作以及最关键的部分——怎么把两者无缝接到一条训练管线里。如果你正在被数据集格式混乱、加载速度慢、数据增强写起来别扭、采样逻辑不灵活这些问题困扰这篇文章应该能帮你把思路理清。1. 两个数据栈先分清它们各自管什么很多人刚接触深度学习数据处理时会被datasets和torch.utils.data搞混觉得都是处理数据到底该用哪个我用了一个很长的阶段才明白这两个东西压根不在同一个层面上它们是上下游关系。1.1 从原始数据到样本库datasets负责的环节Hugging Face的datasets库解决的核心问题是“如何把散落在磁盘上的原始文件变成结构化的、可查询的、可缓存的数据集”。换句话说它的职责是数据整理和清洗。我在实际项目中用到的最典型场景是这样的从某个比赛网站下载了一批图像分类数据文件夹里是几万张图片外加一个标注文件。传统的做法是拿Python脚本自己遍历目录、读CSV、做标签编码、写成一个pickle或者npy文件——这套流程每个项目都得重新写一遍而且一旦数据量超过内存还得手动搞分块、搞缓存相当痛苦。datasets库把这件事抽象成了几个标准动作。load_dataset加载原始数据map做批量转换比如统一图片尺寸、做tokenizationfilter做过滤select做切片train_test_split做划分。每一步都有缓存机制改了转换逻辑只需要重跑变更的那一部分不用从头再来。这个库还有一个很讨喜的设计数据集在磁盘上是以内存映射方式存储的加载一个10GB的数据集并不会真的占满10GB内存而是按需读取。这一点在我处理大规模文本语料和图像数据时帮助极大。1.2 从样本库到训练批次torch.utils.data负责的环节torch.utils.data管的是另一件事训练过程中的样本读取与批次组装。它由几个核心组件构成Dataset负责定义“如何通过索引取出一个样本”DataLoader负责把Dataset里的样本按batch打包、做并行预取、shuffleSampler则决定了取样本的顺序策略。此外还有collate_fn这个容易被忽略但极其重要的函数负责把一个batch里的多个样本拼成张量。这两个库的关系可以这样理解datasets是加工厂负责把原材料做成半成品——干净、统一的样本库torch.utils.data是食堂打饭窗口负责把半成品按份量batch打出来高效地送到训练餐桌上。前者管“数据长什么样”后者管“数据怎么喂给模型”。1.3 为什么两者要配合而不是二选一有人可能会问datasets库本身也能转成PyTorch格式自带to_torch_dataset方法那直接用datasets不就完事了为什么还要单独学torch.utils.data我的回答是datasets的to_torch_dataset本质上只是包装了一个torch的Dataset接口真正决定训练效率的DataLoader行为比如多进程加载、shuffle策略、采样权重、自定义batch拼装全部还是由torch.utils.data来控制的。反过来也一样纯靠自己写Dataset和DataLoader也不是不行大部分教程也都是这么教的。但一旦数据样本量超过几万预处理逻辑复杂需要缓存、格式需要频繁迭代加个字段、清一批脏数据纯手写方案就会变得非常难受。数据打平、tokenize、缓存这些琐碎工作要花掉大量开发时间而这些恰恰是datasets库已经解决得很成熟的问题。所以我的固定组合是用datasets做数据清洗和预处理的基座用torch.utils.data做训练数据管线的控制层。这篇文章后面所有内容都是围绕这个组合展开的。2. datasets库实战从原始文件到高性能样本库这一节我把datasets库最常用的操作串一遍。用的例子是一个典型的图像分类场景本地有一批按类别分文件夹存放的图片还带一个包含额外信息的CSV注释文件。整个过程我逐步操作过很多次这里直接按最顺手的顺序来讲。2.1 加载本地数据load_dataset的两种常见方式处理本地图片分类数据我优先用ImageFolder格式。datasets库自带一个加载器能直接扫描指定root目录下的子文件夹把每个子文件夹名字当作类别标签from datasets import load_dataset, Image # 目录结构假设 # data/ # train/ # cat/ # img_001.jpg # img_002.jpg # dog/ # img_001.jpg # test/ # cat/ # dog/ ds load_dataset(imagefolder, data_dirdata/train, splittrain) print(ds.features)输出里会看到features包含image字段和label字段label是ClassLabel类型已经自动做了类别到整数ID的映射。这套自动映射非常实用省去手动做字典的步骤。如果数据不是按文件夹组织的而是一个目录下所有图片加一个标注文件CSV、JSONL最常见这时用load_dataset(csv)或者load_dataset(json)更合适import pandas as pd df pd.read_csv(annotations.csv) # 假设csv有 image_path 和 label 两列 ds load_dataset(csv, data_filesannotations.csv, splittrain) print(ds[0])但这里有个问题csv里的image_path只是路径字符串并不是真正的图像数据。如果直接拿去训练模型根本没法用。解决办法是map一个函数把路径读成图像from PIL import Image def read_image(example): image_path example[image_path] image Image.open(image_path).convert(RGB) return {image: image} ds ds.map(read_image)读进来的PIL对象在datasets里会被自动包装成Image特征类型之后的处理就统一了。2.2 用map函数做高效批量转换你要理解的背后机制map是datasets库里最核心、也最需要理解透彻的方法。它负责对数据集里的每个样本执行一个转换函数返回新的字段或者修改旧字段。一个最常见的用途是图像标准化和尺寸统一from torchvision import transforms transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) def preprocess(example): image example[image].convert(RGB) example[pixel_values] transform(image) return examplemap带一个batched参数这个参数容易被新手忽略但影响巨大。默认batchedFalse时转换函数是一个个样本调用的如果设为True函数会收到一个包含多个样本的字典列表def preprocess_batched(batch): images [img.convert(RGB) for img in batch[image]] batch[pixel_values] [transform(img) for img in images] return batch ds ds.map(preprocess_batched, batchedTrue)批量处理模式下PIL的批量读图、批量变换可以利用内部优化整体速度通常比逐条调用快一个量级。我在4万张图片的数据集上做过对比batchedTrue的耗时大约是逐条模式的五分之一到六分之一。理解map的时候要注意一点它不只是普通的for循环遍历它会把结果写入磁盘缓存。下次再执行同样的map会直接读缓存不重新计算。所以如果改动了转换函数最好把load_from_cache_fileFalse传给下一次map或者用ds.cleanup_cache_files()清掉无效缓存避免拿到旧版本的转换结果。2.3 过滤、切片与划分整理数据的几个顺手操作filter可以按条件剔除样本。这个操作在清洗脏数据时非常有用# 过滤掉尺寸小于某个阈值的图片 def is_valid(example): w, h example[image].size return w 32 and h 32 ds ds.filter(is_valid)select可以直接按索引切片。注意select支持传入list比如取前1000条做快速原型验证而不需要创建新的数据集文件small_ds ds.select(range(1000))还有train_test_split做划分split_ds ds.train_test_split(test_size0.2, seed42, stratify_by_columnlabel) train_ds split_ds[train] valid_ds split_ds[test]stratify_by_column是datasets里的分层抽样参数能保持划分前后类别比例一致。做分类任务时强烈建议带上这个参数避免随机划分导致某个小类别在训练集里直接消失。这事我在实践中踩过曾有一个稀有类别的样本总量不到50随机划分后训练集里只剩十几个模型在那个类别上几乎学不到东西。2.4 别忽略内存映射和缓存数据量大时的性能关键datasets底层用的是Arrow格式做存储通过内存映射机制按需读取。我在多个项目里的体感是加载几十GB的数据集物理内存占用不会跟着涨上去这一点和直接把整个数据集读进内存的做法有本质区别。但要注意内存映射依赖缓存文件。缓存路径默认在~/.cache/huggingface/datasets如果磁盘空间吃紧或者团队多人共用同一台机器会出现缓存目录爆满的问题。我一般会在每个项目里显式指定缓存位置from datasets import set_cache_dir set_cache_dir(/data/project_cache/datasets)另外有个问题值得提醒如果处理完的数据集包含大量图像张量uint8 TensorArrow压缩率有限缓存膨胀的速度会超出预期。建议在map阶段就把不再需要的大字段丢掉比如原始图像路径、临时中间变量保持缓存精简。3. torch.utils.data核心机制Dataset/DataLoader/Sampler到底在做什么torch.utils.data这套机制我见过太多人只是照抄教程不知道每一块在干什么。结果就是换了个数据集结构不知道Dataset该怎么改训练速度上不去不知道调DataLoader哪个参数类别不均衡不知道Sampler的作用。这一节我把底层逻辑拆开讲透。3.1 Dataset协议一个必须实现__len__和__getitem__的货架从根本上看PyTorch的Dataset只是定义了一个协议通过整数索引能取出样本。它的抽象类长这样from torch.utils.data import Dataset class MyDataset(Dataset): def __len__(self): # 返回样本总数 pass def __getitem__(self, idx): # 返回第idx个样本 pass我的理解是Dataset就是一个带编号的货架DataLoader是取货员。__getitem__决定货架上每个位置放的是什么__len__告诉取货员货架到底有多长。这两个方法的实现质量直接影响训练数据流水线的性能。最容易犯的错误是在__getitem__里做大量重复计算。比如每次都重新读图、重新做resize、重新做toTensor这些操作如果样本数量大会造成巨大的IO开销。正确做法是把一次性预处理放在前面做掉让__getitem__尽量轻量。举一个结合datasets的例子如果我用datasets处理好了数据最偷懒的Dataset是这样class HuggingFaceDatasetWrapper(Dataset): def __init__(self, hf_dataset): self.dataset hf_dataset def __len__(self): return len(self.dataset) def __getitem__(self, idx): return self.dataset[idx]之所以能这么简单是因为datasets的样本本身就是可索引的字典结构我们只需要把idx透传过去就行。DataLoader每次来取货它从货架上拿下来的就是完整样本字典。3.2 DataLoader多进程加载、shuffle、批组装的完整逻辑DataLoader做的事比看上去多得多。它会负责四件事打乱顺序或者按Sampler的策略取序、按batch_size攒成一波、把一波样本通过collate_fn拼成张量、以及通过num_workers开多个子进程并行做上面这些事。一个典型的使用如下from torch.utils.data import DataLoader train_loader DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastFalse )shuffle在DataLoader中的实现并不是对每个epoch重新生成全量索引而是维护一个随机索引排列每个epoch用不同顺序打乱样本顺序。这里有一个我要强调的点如果你的Dataset __getitem__实现很重比如实时做数据增强单纯开shuffle并不能缓解计算压力真正起作用的是num_workers并行加载。num_workers个人建议先从4或者8开始调。worker数太少GPU经常在等数据worker数太多系统会花大量时间做进程间通信和内存拷贝。根据经验在NVMe固态硬盘做图像读取时4到8个worker就能把常见的ResNet级别的训练喂饱如果数据集小且都在内存里num_workers过高反而会因为进程调度开销变慢。pin_memoryTrue在GPU训练时也建议固定打开。它让数据先锁页存放拷贝到GPU显存时会走更快的内存到显存通道。实际体感有时候单靠pin_memory就能让训练速度提升10%以上。3.3 Sampler控制数据顺序的核心不只是shuffle这么简单很多人把Sampler当成DataLoader里一个可有可无的配件这是极大的误解。控制样本取用顺序的就是Samplershuffle只是它的一种具体表现。PyTorch的默认逻辑是如果不传Sampler且shuffleTrueDataLoader内部会构造一个RandomSampler打乱所有索引如果shuffleFalse就按顺序使用SequentialSampler。Sampler的价值在类别不均衡时真正体现出来。处理不平衡数据时最常做的是过采样少类样本或欠采样多类样本这就要用WeightedRandomSampler。它的用法是给每个样本一个权重权重越大被抽中的概率越高from torch.utils.data import WeightedRandomSampler import numpy as np # 为每个样本计算权重 labels np.array([sample[label] for sample in ds]) class_counts np.bincount(labels) class_weights 1.0 / class_counts sample_weights class_weights[labels] sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights), replacementTrue) loader DataLoader(ds, batch_size32, samplersampler)另一个重要的Sampler是DistributedSampler做多卡分布式训练时每个进程拿到的应该是数据的一个分片而不是完整数据集否则每张卡都会重复加载全部数据既不高效还导致每个epoch每张卡看到的样本顺序是乱的。DistributedSampler干的活就是划分索引并打乱。我自己的经验是分清三个概念——Sampler决定取数顺序batch_sampler决定怎么把多个索引组成一个batchcollate_fn决定batch里的样本怎么拼成一个张量。大多数项目只需要调Sampler只有特别复杂的任务才需要动batch_sampler。3.4 collate_fn把样本拼成batch的幕后功臣也是最容易被忽视的默认collate_fn会把一组样本按字段维度堆叠成张量。但如果样本不是规整的Tensor或者样本里有变长序列比如自然语言处理里一个batch内句子长度不同默认collate_fn就会报错。举个例子。我用datasets做文本分类每一条样本是{text: some sentence, label: 0}在DataLoader里如果不指定collate_fn等到要堆叠时text字段是两个长度不一致的字符串pytorch不知道该怎么拼成张量直接抛异常。这时候需要自己写collate_fndef collate_fn(batch): texts [item[text] for item in batch] labels torch.tensor([item[label] for item in batch]) return {texts: texts, labels: labels}如果使用transformers的tokenizer常规做法是在collate阶段把token化并pad到相同长度from transformers import AutoTokenizer tokenizer AutoTokenizer.from_pretrained(bert-base-chinese) def collate_fn(batch): texts [item[text] for item in batch] labels torch.tensor([item[label] for item in batch]) encodings tokenizer(texts, paddingTrue, truncationTrue, max_length128, return_tensorspt) return {input_ids: encodings[input_ids], attention_mask: encodings[attention_mask], labels: labels}这条经验非常关键collate_fn是在子进程里被调用的要保证它不在主进程里依赖重量级对象否则多进程加载时会出现性能下降甚至报错。比如tokenizer实例虽然可以序列化传给子进程但在处理大批量时效率没有直接在主进程里初始化来得好。稳妥方案是让collate_fn内部懒加载tokenizer。4. 桥接datasets与DataLoader三种实战接法及选择思路接下来是整篇文章的操作核心从datasets处理完的样本库到DataLoader训练管线到底怎么接。我这里给出三种接法分别适用于不同场景并且给出我个人的选择建议。4.1 方案一datasets自带to_torch_dataset最省事但不够灵活datasets库内置了to_torch_dataset方法可以把Hugging Face数据集转成一个PyTorch Dataset。这个方案对只需要基础迭代的项目是最快的import torch from datasets import load_dataset ds load_dataset(imagefolder, data_dirdata/train, splittrain) # 假设已经通过map做好了resize和归一化 ds.set_format(typetorch, columns[pixel_values, label]) torch_ds ds.to_torch_dataset() loader DataLoader(torch_ds, batch_size32, shuffleTrue, num_workers4)这里set_format(typetorch)的作用是把字段转成PyTorch张量格式方便后续直接使用。这个方案最大的问题在于定制性弱。如果需要在__getitem__里做在线增强、需要动态返回不同字段或者需要和自定义Sampler配合to_torch_dataset就会显得束手束脚。4.2 方案二自定义Dataset包装Hugging Face数据集兼顾简单与灵活这个方案目前用得最多。它把datasets作为底层存储自己封装一个Dataset子类在__getitem__里做灵活处理from torch.utils.data import Dataset from datasets import load_dataset import torch class HFDataset(Dataset): def __init__(self, hf_ds): self.ds hf_ds def __len__(self): return len(self.ds) def __getitem__(self, idx): sample self.ds[idx] # 这里可以灵活加入在线数据增强、抽样等逻辑 pixel_values torch.tensor(sample[pixel_values]) if isinstance(sample[pixel_values], list) else sample[pixel_values] label torch.tensor(sample[label], dtypetorch.long) return {pixel_values: pixel_values, label: label}这样的好处是底层存储和缓存还是datasets在管内存映射、缓存加速这些优势全部保留上层迭代逻辑完全由自己控制加数据增强、加字段处理都很自由。如果你需要在__getitem__里做数据增强这里有个效率上的建议尽量不要在__getitem__里调用重型变换比如随机裁剪、亮度抖动因为它是逐样本调用的没有batch层面的优化空间。更推荐在datasets的map阶段做一部分确定性预处理在__getitem__阶段只做轻量级在线增强或者干脆用一个额外的transform组件挂在外面统一处理。4.3 方案三让datasets直接输出张量配合自定义collate_fn第三种方案是把数据格式转换放在datasets阶段完成然后让collate_fn负责拼装和增强。这种方式在我做多模态文本图像项目的时候特别好用因为不同字段的预处理方式差异很大。from datasets import load_dataset from torch.utils.data import DataLoader import torch ds load_dataset(json, data_filesmultimodal.jsonl, splittrain) # 假设已经分别对文本和图片字段做了预处理 ds.set_format(typetorch) def multimodal_collate(batch): pixel_values torch.stack([item[pixel_values] for item in batch]) input_ids torch.stack([item[input_ids] for item in batch]) attention_mask torch.stack([item[attention_mask] for item in batch]) labels torch.stack([item[label] for item in batch]) return { pixel_values: pixel_values, input_ids: input_ids, attention_mask: attention_mask, labels: labels } loader DataLoader(ds, batch_size16, collate_fnmultimodal_collate, num_workers4)这个方案里datasets负责把多模态数据变成统一的张量格式collate_fn负责对齐和拼装职责边界清晰。如果后续要改batch拼装方式或者按batch动态padding只需要改collate_fn不需要动数据预处理链路。4.4 三种方案怎么选我的经验判断标准三种方案不是互相替代的关系是适用场景的差异。我的选择标准可以总结成一张表方案适用场景定制能力上手难度性能表现to_torch_dataset快速原型、标准分类任务低最低中自定义Dataset包装需要在线增强、字段动态处理高中高自定义collate_fn多模态、变长序列、复杂拼装最高较高高实际项目里我一般先用方案一把基线跑通模型能出结果了再迭代。一旦涉及数据增强或者特殊的采样逻辑立刻切到方案二。多模态任务直接上方案三。5. 训练管线中的避坑清单与参数调优心得这一节更像是踩坑日志的集合。我把在真实项目里遇到过的、也经常在社区里看到别人问的问题集中整理出来附带我的解决思路。5.1 token化放在map阶段还是collate阶段要分场景做文本分类时token化放哪一直是个争议点。放map阶段预计算好input_ids和attention_mask训练时速度最快放collate阶段每个batch实时tokenize灵活度高但慢。我的建议是分场景数据集固定、不需要频繁修改放map阶段用datasets的缓存一次性算好。训练时直接读张量。我之前处理过20万条中文文本map阶段用tokenizer的parallelism配合batchedTrue大概几分钟就处理完之后每次epoch都不再触碰tokenizer速度非常快。需要在训练过程中动态改变token化参数比如动态mask放collate阶段。但尽量保证tokenizer在collate里是懒加载的不要在创建DataLoader时就实例化否则多进程场景下容易出问题。数据增强包含文本替换扰动放collate阶段。因为map阶段做死的token化无法适应动态扰动。5.2 num_workers不是越大越好瓶颈排查有一个明确方向很多人在训练缓慢时会盲目调大num_workers从4调到16甚至32结果没有变快反而变慢。原因是多进程加载数据有固定的进程创建和IPC进程间通信开销进程数量超过某个阈值后收益会被开销盖过。排查训练缓慢瓶颈的正确顺序是先看GPU利用率。利用率高通常90%以上瓶颈不在数据加载不用动num_workers。利用率低但CPU没占满检查IO。比如图片读取速度是否受限于磁盘随机读性能。在机械硬盘上num_workers再高也快不了换NVMe才有本质提升。利用率低而且CPU已经被占满说明预处理太耗时。此时不是加worker而是要把重计算往map阶段迁移或者换更快的预处理库。我习惯用一个简单脚本测数据管线速度import time from torch.utils.data import DataLoader loader DataLoader(ds, batch_size64, num_workers8, shuffleTrue) start time.time() for i, batch in enumerate(loader): if i 50: break elapsed time.time() - start print(f50 batches in {elapsed:.2f}s, avg {elapsed/50*1000:.1f}ms/batch)拿这个速度对照训练时每个batch的前向反向耗时就能判断数据管线到底够不够快。如果数据管线速度远低于训练耗时说明瓶颈不在数据侧。5.3 shuffle之后还要不要Sampler两者谁优先DataLoader的shuffle参数和自定义Sampler存在互斥如果你传了samplershuffle参数会被忽略。这一点在代码里很容易被忽视导致你以为在shuffle实际上一直按自定义Sampler的顺序取数。我的经验是自定义Sampler时不要依赖DataLoader的shuffle。如果既想按权重采样又要每个epoch打乱顺序请在Sampler内部实现打乱逻辑。一个实用写法import random from torch.utils.data import Sampler class WeightedShuffleSampler(Sampler): def __init__(self, weights, num_samples, replacementTrue): self.weights weights self.num_samples num_samples self.replacement replacement def __iter__(self): if self.replacement: indices list(torch.multinomial(torch.tensor(self.weights, dtypetorch.float32), self.num_samples, replacementTrue).numpy()) else: indices list(range(self.num_samples)) random.shuffle(indices) return iter(indices) def __len__(self): return self.num_samples还有一种常见需求是每个epoch让每个样本恰好出现一次只是在epoch内顺序不同同时类别比例经过重加权。这个逻辑需要自己维护一个索引列表在每个epoch开始时shuffle它。标准WeightedRandomSampler在replacementTrue时无法保证每个样本都被覆盖到它是纯概率抽取有些样本可能多个epoch都不出现。5.4 数据处理幂等性改完转换函数后缓存不更新的坑datasets的缓存是很方便但它的缓存机制基于函数对象的状态和输入数据特征判断是否需要重跑。在实际使用中我遇到过一个典型问题修改了map里的转换函数比如resize尺寸从256改成224但重新运行脚本时datasets可能没有自动检测到函数变更直接用了旧缓存结果训练数据尺寸和模型输入不匹配报错。解决办法是显式关闭缓存读取ds ds.map(my_transform, load_from_cache_fileFalse)或者使用新缓存目录ds ds.map(my_transform, cache_file_name/tmp/new_cache.arrow)养成习惯只要改了转换函数的逻辑就显式指定cache_file_name或者关掉读缓存。千万别迷信自动失效。5.5 图片数据与pin_memory一个容易忽略的显存拷贝细节做图像任务时很多人把数据读成PIL Image然后在__getitem__里转Tensor。这个转换是在CPU上做的数据从CPU内存到GPU显存走的是PCIe通道。如果启用pin_memoryTruePyTorch会把CPU数据放到锁页内存里传输给GPU时可以直接用DMA拷贝省去中间一拷。但要注意如果你在map阶段已经把数据转成了numpy数组或Tensor并通过set_format存成了torch格式这个优化依然有效。真正影响性能的是转Tensor的时机越早转成Tensor训练时越省CPU开销。如果数据集的pixel_values在map阶段就已经是float32的TensorDataLoader加载时直接用torch.stack拼batch此时pin_memory的效率收益非常明显。我做过一次粗略测试同样一个ResNet18训练任务开启pin_memory后每个epoch时间大约缩短8%到15%。6. 一个完整可跑的示例从原始图片到训练循环上面讲了很多原理和细节这一节我直接放一个完整的可运行代码把前面所有的组件串起来。这个示例是一个最小但完整的图像分类训练数据管线你只需要替换数据路径就能跑通。6.1 完整代码datasets清洗自定义Dataset包装DataLoader训练import torch import torch.nn as nn from torch.utils.data import DataLoader, Dataset from torchvision import transforms from datasets import load_dataset import sys # ---------- 1. 加载并预处理数据 ---------- data_dir data/train ds load_dataset(imagefolder, data_dirdata_dir, splittrain) # 划分训练/验证集按类别分层抽样 split_ds ds.train_test_split(test_size0.2, seed42, stratify_by_columnlabel) train_hf split_ds[train] valid_hf split_ds[test] # 定义确定性预处理统一尺寸 转Tensor 归一化 base_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) def preprocess(example): image example[image].convert(RGB) example[pixel_values] base_transform(image) return example # 预处理并缓存注意每次修改preprocess后要设置load_from_cache_fileFalse train_hf train_hf.map(preprocess, batchedFalse, load_from_cache_fileTrue) valid_hf valid_hf.map(preprocess, batchedFalse, load_from_cache_fileTrue) # ---------- 2. 自定义Dataset包装 ---------- class HFDataset(Dataset): def __init__(self, hf_ds, augmentNone): self.ds hf_ds self.augment augment def __len__(self): return len(self.ds) def __getitem__(self, idx): sample self.ds[idx] pixel_values sample[pixel_values] label torch.tensor(sample[label], dtypetorch.long) # 在线增强可选项 if self.augment is not None: pixel_values self.augment(pixel_values) return pixel_values, label # 训练集可以加在线增强 train_augment transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.ColorJitter(brightness0.2, contrast0.2) ]) train_dataset HFDataset(train_hf, augmenttrain_augment) valid_dataset HFDataset(valid_hf) # ---------- 3. 构建DataLoader ---------- train_loader DataLoader( train_dataset, batch_size64, shuffleTrue, num_workers8, pin_memoryTrue ) valid_loader DataLoader( valid_dataset, batch_size64, shuffleFalse, num_workers4, pin_memoryTrue ) # ---------- 4. 定义模型并训练 ---------- model nn.Sequential( nn.Conv2d(3, 16, kernel_size3, padding1), nn.ReLU(), nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(16, train_hf.features[label].num_classes) ) device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() for epoch in range(2): model.train() total_loss 0 for batch_idx, (images, labels) in enumerate(train_loader): images images.to(device, non_blockingTrue) labels labels.to(device, non_blockingTrue) outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() if batch_idx % 10 0: print(fEpoch {epoch} Batch {batch_idx} Loss {loss.item():.4f}) # 验证 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in valid_loader: images images.to(device, non_blockingTrue) labels labels.to(device, non_blockingTrue) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() print(fEpoch {epoch} Val Accuracy: {correct / total:.4f})这段代码看起来长但每一段都是必要的。你只需要把data_dir换成自己的数据路径模型结构换成你的实际模型其余部分可以直接复用。6.2 这个示例里值得注意的四个细节第一train_test_split之后两个子集是基于同一份原始数据划出来的各自会有独立的缓存。如果后续修改了preprocess逻辑两个子集都需要重跑map别只改一个。第二图像在map阶段已经被转成Tensor并存进缓存这意味着DataLoader加载的是Tensor而非PIL对象。此时在线增强作用在Tensor上ColorJitter这类变换对Tensor的格式比较敏感在PyTorch里它接受Tensor输入。如果增强函数要求PIL格式务必在map阶段保留PIL格式或者增强前再做转换。第三代码里to(device, non_blockingTrue)与pin_memoryTrue是配套的。启用pin_memory后non_blockingTrue能让数据传输异步化进一步减少GPU等待。如果pin_memory没开non_blockingTrue不会生效甚至在某些环境下更慢。两者要一起用。第四HFDataset的__getitem__返回的是(pixel_values, label)元组而不是字典。这样写是为了和nn.CrossEntropyLoss的输入风格匹配也方便DataLoader默认collate_fn直接堆叠。如果返回字典需要额外写collate_fn代码会长一些。7. 数据迭代与模型实验数据处理管线也需要版本管理很多初学者把数据处理当成一次性工作跑完map就再也不碰了。但真实项目里数据处理几乎和模型迭代同步演进发现某些类别数据有噪声要重新清洗发现模型在某些样本上表现差要分析数据分布发现数据增强策略过强要调整增强参数。也就是说数据处理管线本身是需要版本管理的。7.1 用datasets的缓存机制建立可回溯的实验记录datasets的缓存机制天然支持这种需求。每跑一次map它会生成一个带哈希名的缓存文件。我们可以把每次map的配置参数记录在实验日志里需要复现某个实验时根据参数重建对应的数据集状态。具体做法是给map函数传一个额外的metadata字段或者把预处理参数存成一个JSON文件连同缓存一起归档。我的个人习惯是transform_config { resize: 224, normalize: imagenet, augment: hflipcolorjitter, date: 2025-01-01 } train_hf train_hf.map( preprocess, fn_kwargs{config: transform_config}, load_from_cache_fileFalse )这样每次实验的transform_config都记录在样本元数据里哪怕过了几个月回来翻旧实验也能快速知道当时用了什么预处理。7.2 数据切分与线下验证的一致性陷阱这个坑我在做比赛时踩过印象非常深训练时用train_test_split随机划分了训练集和验证集提交测试时发现线下验证集分数很高但线上分数很低后来排查发现是数据划分时按整个原始数据集一起划分但原始数据里同一类图片有大量高度相似的重复样本比如连拍帧导致验证集和训练集之间存在信息泄漏。解决办法是要先对样本做去重或者分组再划分。datasets里的add_column配合group操作可以做一些基础的重复样本检测更复杂的情况需要借助对图像做感知哈希。一个简化的做法import hashlib from PIL import Image def compute_image_hash(example): img example[image].convert(RGB) # 简单哈希实际可用感知哈希提高鲁棒性 h hashlib.md5(img.tobytes()).hexdigest() return {hash: h} ds ds.map(compute_image_hash, num_proc4) # 按hash去重 ds ds.unique(hash) # 这个方法仅适用于小数据集大数据集需要group_by或者filter严格来说datasets没有内置group_by去重方法实际项目里我会先把hash列提取到pandas里做去重再把保留的索引传回来select。逻辑不复杂但这一步很关键。7.3 标签噪声处理清洗数据比调模型参数更值得投入我从一个朋友的项目里看到过一个对比同样的模型同样的超参数据清洗前后F1提升了将近8个百分点。这个提升幅度靠调参很难达到。标签噪声最常见的来源是众包标注的模糊边界、数据爬取时引入的错误标签、以及标注规范不统一。处理方式除了人工复核还可以用模型辅助筛选先用一个初步训练好的模型做预测找出预测置信度低或高置信度但预测错误与标签不一致的样本人工过一遍修正标签或者剔除。这个流程在datasets里实现非常顺畅只需要在样本里加一个pred_label字段然后用filter条件筛选。8. 多进程与分布式训练中的数据加载细节最后这一节我专门讲一下数据加载在更复杂环境下的表现。很多项目跑到中途要加速从单卡到多卡或者从单机到多机数据处理管线如果设计得好过渡会很平滑设计得不好各种诡异问题就来了。8.1 num_workers在分布式环境下的正确姿势单机多卡训练时如果每张卡都创建一个DataLoader并且num_workers设置过高总进程数会成倍增长可能造成CPU争用和内存带宽饱和。我建议的做法是每个进程每张卡的num_workers保持和单卡时相同不要因为卡数多了就盲目翻倍。比如单卡时num_workers8四卡时每张卡还是8总计32个worker这时需要用watch命令观察CPU和内存是否吃紧如果吃紧适当降到6或者4。如果做的是多机训练每个机器独立读本地磁盘情况会少很多复杂问题。但如果多个机器共享一个网络文件系统NFS要特别注意并发读的吞吐瓶颈一个常用解法是把数据预处理成多个shard分片文件让每个进程只读一个分片减少并发竞争。8.2 datasets的num_proc参数预处理阶段的并行度datasets的map支持num_proc参数控制预处理并行进程数这和DataLoader的num_workers不是一回事。map的num_proc管的是预处理阶段比如读图、tokenizeDataLoader的num_workers管的是训练阶段的数据迭代。这是两个独立维度的并行都要分别调。我用datasets做预处理时会先确认CPU核心数然后设置num_proc为核心数减一或减二留出主进程的余量。对纯CPU密集的预处理比如图像resize这个并行度通常能把预处理时间压缩到单进程的十分之一。但这个并行在部分数据集上无法使用比如map的函数里依赖了无法序列化的共享对象时需要用num_proc1逐条处理。8.3 数据加载顺序与随机性分布式训练中不可或缺的种子管理分布式训练还有一个容易坑到人的地方不同的进程如果数据顺序不一致模型输出的结果会受影响。需要在每个worker里设置独立的随机种子同时保证每次shuffle的顺序在不同进程间是不同的这样batch和batch之间才有足够的随机性。PyTorch官方推荐的做法是给DataLoader设置generator给每个worker设置worker_init_fnimport random import numpy as np import torch def seed_worker(worker_id): worker_seed torch.initial_seed() % 2**32 random.seed(worker_seed worker_id) np.random.seed(worker_seed worker_id)这个函数会在每个worker进程启动时调用保证不同worker拿到不同的随机序列。如果不设置多个worker可能产生完全相同的增强序列相当于有效batch size减半。关于分布式采样前面提过要用DistributedSampler。还需要记住一点用了DistributedSampler之后DataLoader自带的shuffle参数要设置成False由DistributedSampler内部的shuffle逻辑来控制。两者叠加会有重复打乱的问题不但浪费计算还会导致分布不均匀。我在实际项目中逐步摸索出的这套组合核心就是一句话datasets管数据质量和预处理torch.utils.data管训练效率和迭代逻辑两者边界清晰、配合得当。如果你目前的数据处理链路还是那种脚本读文件手工拼batch一堆全局变量的模式花一个下午把datasets这套流程移植进去我觉得是值得的。数据管线的收益不是即时可见的它会在你每次修改数据处理逻辑、每次跑大规模训练、每次排查速度瓶颈的时候体现出来。数据侧的功夫永远不会白下。
返回列表