深度学习数据处理:datasets与torch.utils.data实战指南
·
在深度学习和数据处理中,datasets(通常指 Hugging Face 的 datasets 库)和 torch.utils.data(PyTorch 内置的数据加载模块)是两个核心工具包。它们分别用于高效加载/处理数据集和构建可迭代的数据管道。下面系统梳理它们的主要类和方法。
一、Hugging Face datasets 库常用方法
安装:
pip install datasets
1. 核心类:Dataset 和 DatasetDict
Dataset:表示单个数据集(类似表格,每列是一个特征)。DatasetDict:包含多个子集(如{"train": ..., "validation": ..., "test": ...})。
2. 常用方法
| 方法 | 说明 | 示例 |
|---|---|---|
load_dataset(path, ...) | 加载内置或自定义数据集 | load_dataset("glue", "mrpc") |
Dataset.map(function, ...) | 最常用! 对每个样本应用函数(支持多进程、缓存) | ds.map(lambda x: {"len": len(x["text"])}) |
Dataset.filter(function) | 过滤样本 | ds.filter(lambda x: len(x["text"]) > 10) |
Dataset.select(indices) | 按索引选择子集 | ds.select([0, 1, 2]) |
Dataset.shuffle(seed=42) | 打乱数据 | ds.shuffle() |
Dataset.train_test_split(...) | 划分训练/测试集 | ds.train_test_split(test_size=0.2) |
Dataset.rename_column(old, new) | 重命名列 | ds.rename_column("label", "class") |
Dataset.remove_columns(names) | 删除列 | ds.remove_columns(["id"]) |
Dataset.to_pandas() | 转为 Pandas DataFrame | df = ds.to_pandas() |
Dataset.from_dict(dict) | 从字典创建 Dataset | Dataset.from_dict({"text": ["a", "b"]}) |
Dataset.save_to_disk(path) | 保存到磁盘(高效格式) | ds.save_to_disk("./my_ds") |
load_from_disk(path) | 从磁盘加载 | ds = load_from_disk("./my_ds") |
3. 特殊功能
- 流式加载(Streaming):处理超大数据集
ds = load_dataset("c4", "en", split="train", streaming=True) - 与 PyTorch/TensorFlow 集成:
ds.set_format(type="torch", columns=["input_ids", "labels"])
二、torch.utils.data 常用组件
属于 PyTorch 核心库,无需额外安装。
1. 核心抽象类
| 类 | 作用 |
|---|---|
Dataset | 抽象类,用户需继承并实现 __len__ 和 __getitem__ |
IterableDataset | 用于流式/无限数据源(如网络数据流) |
DataLoader | 最重要! 将 Dataset 转为可迭代的 batch 数据流 |
2. Dataset 子类(内置实用类)
| 类 | 说明 |
|---|---|
TensorDataset | 从张量直接构建数据集(适用于小数据) |
ConcatDataset | 拼接多个数据集 |
Subset | 取数据集子集(常用于划分训练/验证) |
ChainDataset | 顺序遍历多个数据集(不拼接) |
示例:TensorDataset
from torch.utils.data import TensorDataset, DataLoader
import torch
X = torch.randn(100, 10)
y = torch.randint(0, 2, (100,))
dataset = TensorDataset(X, y)
loader = DataLoader(dataset, batch_size=16, shuffle=True)
3. DataLoader 关键参数
| 参数 | 说明 |
|---|---|
dataset | 输入的 Dataset 实例 |
batch_size | 批大小 |
shuffle | 是否打乱(仅对 map-style dataset 有效) |
num_workers | 多进程加载(>0 可加速 I/O) |
collate_fn | 重要! 自定义如何将样本列表拼成 batch(处理变长序列必备) |
pin_memory | GPU 训练时设为 True 可提速 |
drop_last | 最后一个 batch 不足 batch_size 时是否丢弃 |
自定义 collate_fn 示例(处理文本长度不一):
def collate_fn(batch):
texts = [item["text"] for item in batch]
labels = torch.tensor([item["label"] for item in batch])
# 使用 tokenizer padding
encodings = tokenizer(texts, padding=True, return_tensors="pt")
return {**encodings, "labels": labels}
loader = DataLoader(dataset, batch_size=8, collate_fn=collate_fn)
4. 工具函数
| 函数 | 说明 |
|---|---|
random_split(dataset, lengths) | 随机划分数据集(替代 Subset + 手动索引) |
get_worker_info() | 在 Dataset 中获取当前 worker 信息(用于多进程) |
示例:划分训练/验证集
train_ds, val_ds = random_split(dataset, [50000, 10000])
三、典型工作流整合(Hugging Face + PyTorch)
from datasets import load_dataset
from torch.utils.data import DataLoader
from transformers import AutoTokenizer
# 1. 加载数据集
ds = load_dataset("imdb")
# 2. 预处理(使用 map)
tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")
def tokenize(examples):
return tokenizer(examples["text"], truncation=True, padding="max_length", max_length=512)
ds = ds.map(tokenize, batched=True)
ds.set_format(type="torch", columns=["input_ids", "attention_mask", "label"])
# 3. 构建 DataLoader
train_loader = DataLoader(ds["train"], batch_size=16, shuffle=True)
val_loader = DataLoader(ds["test"], batch_size=16)
✅ 此流程中:
datasets.map()完成高效预处理(可缓存、多进程)set_format(type="torch")自动转为 PyTorch 张量DataLoader负责 batch、shuffle、多进程加载
四、总结对比
| 功能 | datasets | torch.utils.data |
|---|---|---|
| 数据加载 | ✅ 支持本地/远程/流式 | ❌ 需自行实现 Dataset |
| 数据预处理 | ✅ map, filter 等高级操作 | ❌ 通常在 __getitem__ 中简单处理 |
| 内存效率 | ✅ 惰性计算 + 磁盘缓存 | ⚠️ 全加载到内存(除非自定义流式) |
| 与 PyTorch 集成 | ✅ 通过 set_format | ✅ 原生支持 |
| 适用阶段 | 数据准备 & 预处理 | 数据加载 & batch 生成 |
💡 最佳实践:
用datasets做预处理 → 用DataLoader做训练加载,两者互补,构成现代 NLP/ML 标准 pipeline。
如需具体场景示例(如“如何用 IterableDataset 处理超大文件”或“自定义 Dataset 读取图像”),欢迎继续提问!
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)