资讯动态

深度学习数据处理流水线:datasets 与 torch.utils.data 实战指南

发布时间:2026/9/30 11:48:42 来源:尧图企业网站定制
做深度学习项目我见过太多人把精力全砸在模型结构上等到开始写训练代码才发现数据从原始 csv 变成模型能吃的 mini-batch中间的坑比模型本身还多。这篇实战指南要解决的就是这一整段“深度学习数据处理”流水线用 Hugging Face 的 datasets 库把脏乱的原始数据整理成干净、可复现的结构化数据用 PyTorch 自带的 torch.utils.data 把数据按 batch 稳定地喂给模型。两个工具配合好了清洗、切分、缓存、打乱、并行加载这些常规操作都有标准解法不用每次项目都临时写一堆 for 循环和全局变量去凑合。文章面向正在入门深度学习的同学也适合那些已经被 DataLoader 报错折磨过的实战选手目标是让你看完就能直接照搬这套流程。1. 先搞清楚两个库各自该干什么datasets 和 torch.utils.data 听起来都带“data”但它们是两个完全不同的层面。一个是数据处理框架一个是模型训练的数据供给器混在一起用没问题但如果你不知道它们各自解决什么问题很容易写出又慢又难维护的数据代码。1.1 一条深度学习数据流水线的两段分工我习惯把数据流拆成“加工车间”和“传送带”两个角色。datasets 是加工车间它负责把原始数据读进来、清洗、过滤、切分、缓存中间结果。比如你有一份 20 万条的新闻标题 csv里面有缺失值、乱码、重复项还夹杂着空白行这些脏活累活都是 datasets 的职责。它处理完之后数据应该是一张结构清晰、每一列都有明确含义的“表”。torch.utils.data 是传送带它不管数据干不干净只管怎么把已经处理好的数据一条条、一捆捆地送到训练循环里。DataLoader 负责按 batch_size 打包、按 shuffle 打乱顺序、按 num_workers 多进程预取。它追求的是“训练循环拿到手的永远是一个形状正确、类型正确的 Tensor”至于这个 Tensor 是从内存里还是磁盘缓存里来的它不关心。这个分工很重要。很多人一上来就想着自己写 Dataset 子类在getitem里做清洗、过滤、去重结果每个 epoch 都要重复执行一遍清洗逻辑训练速度慢得离谱。正确做法是清洗和预处理在 datasets 阶段只做一次之后 DataLoader 只负责高效索引和打包。1.2 为什么不直接拿 dict 和 list 硬扛有朋友问过我数据处理这么简单直接读 csv 到一个 list of dict然后循环里自己切片不就行了小数据集确实可以但一旦数据量上来问题就暴露了。第一是内存。20 万条文本每条 500 字光 Python 字符串对象和 dict 的包装开销就能吃掉好几个 GB。datasets 底层用的是 Arrow 列式存储数值列是连续内存的数组字符串列也有高效的编码方式内存占用通常只有纯 Python 结构的五分之一到三分之一。第二是缓存。datasets 的 map 操作会缓存中间结果第一次清洗完写入磁盘第二次不管你在哪个脚本里加载只要缓存 key 没变就直接命中不用重新跑一遍。dict/list 做不到这一点每次跑脚本都得从头清洗。第三是并行。datasets 的 map 可以开 num_proc 多进程处理几行代码就能把 CPU 核数吃满手写 for 循环想达到同样的效果你得自己搞 multiprocessing麻烦且容易出错。所以我的结论是数据量在几万条以下怎么折腾都行数据量上来之后datasets 的缓存和列式存储带来的收益是决定性的。既然深度学习基本没有小数据量场景那从一开始就用标准工具就是最省事的选择。2. datasets 库的实战操作加载、清洗、切分datasets 库的功能非常多但真正高频用到的就几个load_dataset、map、filter、select、shuffle、train_test_split还有 set_format / with_format。把这几个吃透已经能覆盖绝大多数深度学习数据处理场景。2.1 从 load_dataset 开始本地 csv 和 Hub 数据都吃load_dataset 是入口函数它既能加载 Hugging Face Hub 上的公开数据集也能加载本地文件。加载本地 csv 的写法很直接from datasets import load_dataset # 加载单个 csv ds load_dataset(csv, data_filestrain.csv, splittrain) # 加载多个 csv 并合并 ds load_dataset(csv, data_files[train_part1.csv, train_part2.csv], splittrain)这里有个常用参数 delimiter如果你的文件是制表符分隔或者分号分隔要显式传进去。还有一个容易被忽略的是一致性如果你后续要加载验证集最好在同一个 load_dataset 调用里同时加载dataset load_dataset( csv, data_files{train: train.csv, test: test.csv} )这样返回的是一个 DatasetDict包含 train 和 test 两个子集后面做映射、过滤时非常方便。DatasetDict 的操作接口和单个 Dataset 几乎一样可以对 train、test 分别操作也可以直接用 dataset.map(...) 同时作用于所有子集。如果只是测试流程不想真的下载大文件可以用 datasets 里的示例数据快速验证ds load_dataset(imdb, splittrain[:100])split 切片语法是 datasets 的一大特色train[:100] 表示只取前 100 条train[:80%] 表示取前 80%。这个语法在调试 pipeline 时特别实用不需要真的把全量数据加载进来。2.2 map 是真正的主角清洗逻辑都在这map 是 datasets 最核心的方法所有“每条样本做一个变换”的操作都靠它。它的基本逻辑是对每一条数据执行一个函数把函数返回的字典里的新键值合并进原样本。看个具体例子def clean_text(example): text example[text].strip() return {clean_text: text, seq_len: len(text)} ds ds.map(clean_text, num_proc4)这里函数接收 example 字典返回一个新字典。返回的键如果是原来的列就覆盖如果是新键就新增一列。所以执行完之后ds 里会多出 clean_text 和 seq_len 两列原始 text 列还在。这个“保留原数据、新增处理结果”的设计非常贴心调试的时候可以随时对比原始值和清洗值。map 有几个参数是实战中必须掌握的。第一个是 num_proc并行进程数。清洗之类的轻量操作开 4 到 8 个进程一般能把 CPU 占满。但要注意 Windows 下 num_proc 有多进程坑后面第五节详细说。第二个是 batched 参数默认 False 表示逐条调用函数设成 True 时函数接收的是一个列表可以批量处理。典型场景是分词BERT 这类 tokenizer 本身支持 batch 编码用 batchedTrue 速度能提升好几倍。写法是这样的def tokenize_batch(examples): return tokenizer( examples[clean_text], truncationTrue, max_length128, return_attention_maskFalse, ) ds ds.map(tokenize_batch, batchedTrue, num_proc4)注意 batchedTrue 时函数返回的字典里每个 value 应该是一个列表列表长度等于这个 batch 的样本数。第三个是 remove_columns如果你确定原始列没用了可以在 map 时顺手删掉省内存ds ds.map(tokenize_batch, batchedTrue, remove_columns[text])最后要强调一个缓存机制map 的结果会缓存到磁盘上。好处是第二次跑同样的 map 直接读缓存秒开。坏处是你改了清洗函数内部逻辑但没改调用参数缓存可能不会自动失效导致你看到的还是旧结果。这个坑在 5.3 节详细说。2.3 filter、select、shuffle、train_test_split四个最容易被搞混的操作这四个操作都是“返回新 Dataset”不会修改原对象理解这一点你就不会把数据集搞乱。filter 按条件筛掉不需要的样本函数返回 True 则保留。比如我只想要标签是 0 和 1 的样本ds ds.filter(lambda x: x[label] in [0, 1])select 按索引位置挑样本作用是重排或只保留指定索引。比如我想看看前 100 条长什么样可以 ds.select(range(100))。它和 filter 的区别在于filter 是“按值筛选”select 是“按位置筛选”。如果你已经知道要保留哪些行的索引select 比 filter 快得多因为它不做条件判断直接按位置取。shuffle 打乱顺序。这个操作在划分训练验证集之前特别容易被误用。很多人喜欢先 shuffle 再从前面切一刀来划分数据这不是不行但 datasets 内置的 train_test_split 更规范。split_ds ds.train_test_split(test_size0.2, seed42) # split_ds 是 DatasetDict含 train 和 testtrain_test_split 有几个实用参数test_size 可以是比例也可以是绝对条数seed 控制随机性保证两次划分结果一致stratify_by_column 可以按某列分层抽样适合分类不平衡的数据集。这些操作可以链式调用读起来很清楚result ( ds .filter(lambda x: x[label] in [0, 1]) .shuffle(seed42) .train_test_split(test_size0.2, seed42) )2.4 set_format 与 with_format让数据直接成了 Tensordatasets 的最高频使用场景之一就是和 PyTorch 无缝衔接。默认情况下你索引一条数据拿到的是 Python 原生类型比如 input_ids 是 list、label 是 int。但是用 with_format 可以指定返回类型为 torchds ds.with_format(torch, columns[input_ids, attention_mask, label]) sample ds[0] # sample[input_ids] 已经是一个 torch.Tensorset_format 和 with_format 的区别在于set_format 是原地修改with_format 是返回一个新对象不污染原对象。我建议优先用 with_format尤其是调试多个格式切换的时候原对象保持不变更安全。要注意的是with_format(torch) 只是把列转成了 Tensor并不会自动把多个列打包成一个 (input, label) 元组。DataLoader 在 batch 的时候会默认把每条样本堆叠成 dict of tensors这在自定义模型里用起来不够直接。所以实践中更常见的是配合第三部分的 collate_fn 来重新组织输出格式。3. torch.utils.data 的正确打开方式现在数据已经清洗干净、切分好了接下来就是让 DataLoader 把这些数据高效地喂给模型。这一部分的核心有两个什么时候需要自己写 Dataset 子类以及 DataLoader 参数到底怎么调。3.1 什么时候要自己写 Dataset 子类很多人一提起 torch.utils.data 就想到继承 Dataset 然后实现len和getitem。但我得说实话你要是在 datasets 阶段已经把数据整理好了90% 的情况下根本不需要自己写 Dataset 子类。为什么因为 datasets.Dataset 本身就实现了len和getitem它天然可以当 torch.utils.data.Dataset 用。你只需要用 with_format(torch) 把数据列转成 Tensor然后直接传给 DataLoader 就行了。自己写子类的场景主要集中在两类一是数据格式非常特殊比如要从数据库实时查询、要做图像在线增强、要读取巨大文件的一部分二是原生数据的返回结构不是你想要的必须做一次重映射。如果你确实要写那核心就是这两个方法from torch.utils.data import Dataset class MyDataset(Dataset): def __init__(self, ds): self.ds ds def __len__(self): return len(self.ds) def __getitem__(self, idx): item self.ds[idx] return item[input_ids], item[label]看起来很简单但有几个细节值得说。第一getitem里尽量只做索引和轻量转换不要放复杂逻辑因为 DataLoader 在多进程模式下会反复调用它逻辑越重开销越大。第二返回结构要统一要么全部返回元组要么全部返回 dict不要有时返回 dict 有时返回 tuple否则后面 collate_fn 会崩。第三如果数据量特别大可以考虑把切分好后的数据导出成 parquet 或内存映射格式避免每次启动都重新加载。3.2 DataLoader 核心参数逐个说清楚DataLoader 是 torch.utils.data 的灵魂它的参数你真的都弄明白了吗我见过太多人在 num_workers、pin_memory、drop_last 这几个参数上凭感觉乱设导致训练速度上不去或者每个 epoch 最后一步形状对不齐。先看最常用的几个参数我整理成了一张表。参数作用建议值备注batch_size每个 batch 的样本数根据 GPU 显存和模型大小调过小训练不稳过大显存爆炸shuffle每个 epoch 是否打乱训练集 True验证集 False验证集打乱没有意义num_workers加载数据的子进程数Windows 建议 0 或 2Linux 可 4~8不是越大越好太大反而拖慢drop_last最后一批不够 batch_size 时是否丢弃训练集建议 True避免 batch 形状不一致pin_memory是否锁页内存加速 GPU 传输GPU 训练建议 True配合 CUDA 使用persistent_workers子进程是否跨 epoch 常驻epoch 多且 num_workers0 时建议 True减少反复创建进程的开销我重点说两个大家容易忽视的。第一个是 num_workers。它并不是越大越好。每个 worker 都是独立 Python 进程它们需要从主进程拷贝数据如果数据量不大或者操作很快进程间通信的代价会大于并行带来的收益。我实测过在 Linux 上文本分类这种轻量数据num_workers4 和 num_workers8 差别不大再往上甚至会变慢。在 Windows 上多进程 worker 还有一个臭名昭著的坑必须在 ifname main 保护下创建 DataLoader否则会不断递归启动新进程5.1 节细说。第二个是 pin_memory。这个参数很多人不理解我解释得通俗一点GPU 要数据得先从内存拷贝过去pin_memory 就是给这块内存加了个“专属锁”让 GPU 直接通过 DMA 搬数据省掉了中间环节。如果你的数据已经加载到了 CPU 内存并且是 GPU 训练建议无脑设成 True。但它不是免费的锁页内存比较稀缺开太多会影响系统其他部分一般默认 True 即可。3.3 从 datasets 到 DataLoader 的三种衔接路径到底怎么把 datasets.Dataset 喂给 DataLoader我梳理了三种常见路径按推荐程度排序。路径一直接传给 DataLoader。这是最偷懒也最常用的一种loader DataLoader( ds.with_format(torch), batch_size32, shuffleTrue, num_workers4, )这样做可行但 DataLoader 返回的每个 batch 是一个 dict里面是 input_ids、attention_mask、label 这些列名对应的 Tensor。如果你的模型 forward 方法接收多个命名参数这种结构其实挺方便。唯一的问题是很多人习惯 for x, y in loader 这种解构写法用 dict 就要改成 for batch in loader 然后 batch[input_ids]需要适应。路径二包装一层自定义 Dataset 子类把返回结构改成 (input, label)。这种最符合大多数人的 PyTorch 习惯class TorchDataset(Dataset): def __init__(self, ds): self.ds ds def __len__(self): return len(self.ds) def __getitem__(self, idx): item self.ds[idx] return item[input_ids], item[label]路径三直接把 datasets.Dataset 切片成 list 再喂给 TensorDataset。这个适合小数据量快速测试方便但内存开销大不推荐上生产。我个人最喜欢的组合是datasets 阶段做清洗和预处理然后用 with_format(torch) 转换列最后用路径二的三行包装类把结构改成元组。这样既保留了 datasets 的缓存优势又不别扭地改变 PyTorch 的使用习惯。3.4 collate_fn 是变长数据的救命稻草如果你处理的是文本、变长序列、或者不同尺寸的图就一定会遇到“一个 batch 里的样本形状不一致”的问题。DataLoader 默认的 collate 逻辑是把一堆样本沿着第 0 维堆叠它要求每个样本的形状是相同的。遇到变长文本就会报错说 sizes must match。解决办法是自定义 collate_fn。它的输入是一个长度为 batch_size 的 list每个元素是getitem返回的一个样本输出是你要喂给模型的 Tensor。比如处理不定长文本时最常用的策略是 pad 到 batch 内最大长度def collate_fn(batch): input_ids [torch.tensor(item[0]) for item in batch] labels torch.tensor([item[1] for item in batch]) max_len max(t.size(0) for t in input_ids) padded torch.zeros(len(batch), max_len, dtypetorch.long) for i, t in enumerate(input_ids): padded[i, :t.size(0)] t return padded, labelscollate_fn 里可以做任何复杂的 batch 拼接逻辑padding、按长度排序、生成 attention_mask、甚至做 batch 级别的数据增强。写法也灵活可以是函数也可以是类。新手常犯的错误是在getitem里把样本 pad 成全局最大长度。这样确实能保证形状一致但代价是内存浪费和训练变慢。正确做法是在 collate_fn 里 pad 到 batch 内最大长度而不是全局最大长度。一个 batch 内部形状一致就够了模型接收的是 batch 数据不是全局数据。4. 端到端实战从 20 万条原始文本到可训练 DataLoader理论说了这么多配一个完整流程才有感觉。我拿一个真实的文本分类场景走一遍全流程从原始 csv 到 DataLoader 一步不缺。4.1 项目背景与数据长什么样假设我们要做一个新闻标题分类模型任务是判断一条新闻标题属于“科技、体育、财经”中的哪一类。原始数据是一份 20 万条的 csv字段就两列title 是标题文本label 是字符串类型的类别标签。这个数据存在的问题很典型部分 title 为空或全空格、有重复标题、label 有少量拼写变体比如“Tech”和“technology”混在一起。这不是什么极端情况几乎每个真实项目都会遇到这类脏数据。我先定义一个标签映射让所有处理逻辑都从字符串类别映射到数字 idlabel2id {科技: 0, 体育: 1, 财经: 2} id2label {v: k for k, v in label2id.items()}这一步看起来简单但非常重要。模型只能吃数字你还需要保留 id2label 做预测结果转换。我建议在项目一开始就定义好这个映射并固定下来不要在处理过程中临时改。4.2 完整数据处理流水线第一步是加载数据。使用 load_dataset 加载本地 csv这一步会把文件读成一个 Arrow 表格内存占用比纯 Python 小得多。from datasets import load_dataset ds load_dataset(csv, data_filesnews_titles.csv, splittrain) print(ds) print(ds[0])打印出来你会看到每一条是一个 dict包含 title 和 label 两列。第二步是清洗。我定义一个清洗函数去掉标题首尾空格、过滤空标题、把 label 统一成小写、过滤掉不在映射里的类别顺便把类别转换成 id。def clean_and_label(example): title example[title].strip() label example[label].strip().lower() label_id label2id.get(label) return {clean_title: title, label_id: label_id} ds ds.map(clean_and_label, num_proc4) ds ds.filter(lambda x: len(x[clean_title]) 0 and x[label_id] is not None)这里有个细节值得说我故意没有在一开始就删掉原始 title 列而是新增 clean_title。这样如果后面发现清洗逻辑有问题还能回头对照原始值。调试期保留原始列确认无误后再用 remove_columns 清理这是我多年养成的习惯。第三步是去重。标题类数据经常有重复重复样本会放大某些类别的占比影响训练。datasets 没有内置去重方法但可以用一个简单技巧把标题哈希成一条 ID然后按 ID 去重。from datasets import Dataset import hashlib def hash_title(example): return {hash: hashlib.md5(example[clean_title].encode()).hexdigest()} ds ds.map(hash_title, num_proc4) def dedup(ds, keyhash): seen set() indices [] for i, h in enumerate(ds[key]): if h not in seen: seen.add(h) indices.append(i) return ds.select(indices) ds dedup(ds)这个写法虽然有一点 O(n) 的 Python 循环但胜在简单直观。数据量大到连这个循环都跑不动时可以考虑用 pandas 先做 groupby 去重再导回但大多数文本分类场景这个量级完全够用。第四步是分词。分词器这里我用 Hugging Face 生态的 tokenizer它的输出可以直接作为模型输入。注意用 batchedTrue 批量处理。from transformers import AutoTokenizer tokenizer AutoTokenizer.from_pretrained(bert-base-chinese) def tokenize_batch(examples): return tokenizer( examples[clean_title], truncationTrue, max_length64, ) ds ds.map(tokenize_batch, batchedTrue, num_proc4, remove_columns[clean_title])这里我截断到了 64。新闻标题一般都不长这个长度既能覆盖绝大多数样本又不会浪费算力。如果你的文本很长一定要根据你的数据分布统计长度分布而不是拍脑袋选个 512。第五步是划分数据集。用 train_test_split 切顺便确认类别比例。split_ds ds.train_test_split(test_size0.2, seed42, stratify_by_columnlabel_id) train_ds split_ds[train] valid_ds split_ds[test]stratify_by_column 会按 label_id 做分层抽样保证训练和验证集的类别比例一致。分类问题强烈建议开这个参数否则遇到过某个类别样本很少、划分后验证集里几乎没有这类样本模型学了个寂寞。第六步是转换成 torch 格式并构建 DataLoader。这里我用一个极简的包装类让 sample 返回(input_ids, attention_mask, label_id)三元组。from torch.utils.data import DataLoader, Dataset class HFDataset(Dataset): def __init__(self, hf_ds): self.ds hf_ds.with_format(torch) def __len__(self): return len(self.ds) def __getitem__(self, idx): item self.ds[idx] return item[input_ids], item[attention_mask], item[label_id] def build_loader(hf_ds, batch_size32, shuffleTrue): return DataLoader( HFDataset(hf_ds), batch_sizebatch_size, shuffleshuffle, num_workers4, drop_lastshuffle, pin_memoryTrue, ) train_loader build_loader(train_ds, batch_size32, shuffleTrue) valid_loader build_loader(valid_ds, batch_size32, shuffleFalse)由于分词后每条样本都截断到了 64 的长度每个 batch 大小都一致这里不需要自定义 collate_fn。如果你的任务里有变长序列且不截断那就必须像 3.4 节那样写一个 collate_fn 做 batch 内 padding。4.3 训练循环里的最后校验DataLoader 搭好之后别急着开训先在训练循环外面做一轮批次校验确认形状和值都符合预期。这一步能帮你把大部分问题挡在训练之前省得训到一半才发现数据错了。for batch in train_loader: input_ids, attention_mask, labels batch print(input_ids.shape, attention_mask.shape, labels.shape) print(labels.unique()) break我实测过这个简单的检查能暴露三类高频问题一是 shape 不对比如忘了包成 Tensor二是 labels 里有异常值比如出现了 3 而你的分类头只有 3 个输出三是批次顺序没打乱比如相邻 batch 的 label 完全一样。这些问题在训练前用两行代码暴露出来比训练时看 loss 炸掉再回头排查省时省力得多。做完这一步数据流水线就可以正式投入训练了。5. 常见问题与排查技巧实录最后这部分是我这些年踩过的坑汇总。每一个都是真实遇到、真实解决过的希望能帮你少走弯路。5.1 num_workers 在 Windows 下疯狂报错这个问题在 Windows 上几乎人人都会遇到DataLoader 设置了 num_workers 大于 0一跑就报 RuntimeError提示 An attempt has been made to start a new process before the current process has finished its bootstrapping phase。原因很简单Windows 没有 Linux 的 fork 机制创建子进程时会重新导入主模块。如果你的脚本顶层就写了创建 DataLoader 的代码子进程一启动就会又执行一遍这段代码于是无限递归。解决的办法也简单把 DataLoader 的创建和训练循环放到 ifname main: 保护的代码块里。如果你的数据加载代码写在函数里那就没问题因为函数在导入时不会执行。我自己的习惯是数据处理脚本统一写成函数调用式主入口就一行if __name__ __main__: run_train()5.2 map 用了 lambda 导致多进程 pickle 失败datasets 的 map 开 num_proc 时会把处理函数通过 pickle 序列化传给子进程。lambda 函数是无法 pickle 的所以你会看到 PicklingError 或者 AttributeError: Cant pickle local object。解决办法很简单把 lambda 改成具名函数定义在模块顶层。比如# 错误写法 ds ds.map(lambda x: {seq_len: len(x[text])}, num_proc4) # 正确写法 def add_seq_len(example): return {seq_len: len(example[text])} ds ds.map(add_seq_len, num_proc4)这个坑在调试时特别烦因为单进程你不开 num_proc 完全正常一开多进程就炸。我的建议是凡是传给 map 或 filter 的函数一律在模块顶层用 def 定义不要图方便用 lambda。5.3 改完预处理代码结果还是老样子datasets 的缓存机制很强大但也会骗人。你明明改了 map 里的清洗逻辑重新跑脚本结果还是旧数据。原因是缓存 key 只跟数据集版本、map 参数、函数源码的 hash 有关有时候你改了函数内部调用的另一个函数这个 hash 变化并没有被正确捕获。验证很简单看打印的信息。如果跑 map 时输出显示 Loading cached processed dataset而不是 Running map说明走的是缓存。强制绕过缓存的方法是指定重新加载ds ds.map(clean_text, num_proc4, load_from_cache_fileFalse)这个参数会让它忽略缓存重新计算。我调试预处理逻辑时习惯先加 load_from_cache_fileFalse 跑一次确认结果对了再撤掉参数依赖缓存。这样既不会踩缓存欺骗也不会为了调试浪费重复计算。5.4 标签对不齐字符串标签映射错了标签对不齐是分类任务最常见的低级错误。症状是训练时 loss 在下降但准确率一直很低或者验证集上某些类别永远预测不对。排查思路就一条确认你的字符串标签到数字 id 的映射和模型输出的 logits 维度一一对应。我自己的习惯是在处理管线里增加一个标签分布打印from collections import Counter print(Counter(ds[label_id]))至少打印一次训练集和验证集的标签分布确认没有异常的类别。另外如果你的 label2id 是通过字符串排序自动生成的比如 {label: i for i, label in enumerate(sorted(set(labels)))}那一定要保证训练和验证用同一个映射字典不要各自生成一份。否则训练集里“科技”是 0、验证集里“科技”成了 2模型永远学不会。5.5 内存爆炸千万别把所有数据 list 化有些人会把 datasets.Dataset 转换成 Python list图的是“方便”。数据量小时没问题20 万条文本转成 list一下子吃掉好几个 GB 内存机器直接卡死。datasets 的优势在于底层是 Arrow 列式存储它是按列存的、内存紧凑不需要像 Python 的 list of dict 那样每一条都带一个哈希表结构。很多在 list 上 write-heavy 的操作在 datasets 上应该用 map 或 filter 的标准方式完成。如果你真的需要在内存里快速遍历所有样本也不要用 list 硬存所有内容而是用生成器逐条处理for i in range(len(ds)): item ds[i] # 逐条处理不占额外内存5.6 数据泄漏划分前做全局清洗是重灾区数据泄漏在数据处理阶段不明显但影响非常大。最常见的泄漏来源是在 train_test_split 之前对整个数据集做了全局的标准化、归一化或者做了全量去重。这会让模型在训练时已经“见过”验证集的信息导致验证指标虚高上线后暴跌。文本分类里一个具体场景是你在全量数据上用 TF-IDF 或者词频统计做特征然后再划分训练验证集。这时候验证集的词表已经是从全量数据里学来的和训练集不再独立。正确做法是先划分再在各子集上单独做统计类预处理。datasets 的 map 本身是按样本独立操作不会造成泄漏但一旦涉及全局统计量就必须先切分再统计这一点务必养成习惯。我在实际项目中还遇到过更隐蔽的泄漏对验证集也做了和训练集完全相同的去重和清洗。如果清洗规则里有“删除重复数据”而且重复数据恰好横跨训练和验证集那就等于同一条内容出现在了两边。数据划分前只能做无关痛痒的字段裁剪所有跟数据内容相关的变换都要放到划分之后。写在最后我始终觉得深度学习项目里最难搞的不是模型结构调参而是把数据处理这条流水线做得又快又稳。datasets 和 torch.utils.data 这对组合一个负责把数据收拾干净、缓存好一个负责高效、稳定地把数据喂给模型配合起来基本上能覆盖 90% 的数据处理需求。最后再分享一个小技巧处理复杂数据集时把每一步 map 的结果都用新列保存比如 raw_text、clean_text、tokenized不要覆盖原始列。这样当模型效果不对劲时你可以逐列排查看看到底是清洗逻辑错了、分词参数错了还是标签映射错了。我靠着这个习惯省下了无数排查时间。数据处理看起来是脏活累活但用对工具、建好流水线它反而会变成整个项目里最省心的一环。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑