在深度学习和数据处理中,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 DataFramedf = ds.to_pandas()
Dataset.from_dict(dict)从字典创建 DatasetDataset.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_memoryGPU 训练时设为 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、多进程加载

四、总结对比

功能datasetstorch.utils.data
数据加载✅ 支持本地/远程/流式❌ 需自行实现 Dataset
数据预处理✅ map, filter 等高级操作❌ 通常在 __getitem__ 中简单处理
内存效率✅ 惰性计算 + 磁盘缓存⚠️ 全加载到内存(除非自定义流式)
与 PyTorch 集成✅ 通过 set_format✅ 原生支持
适用阶段数据准备 & 预处理数据加载 & batch 生成

💡 最佳实践:
用 datasets 做预处理 → 用 DataLoader 做训练加载,两者互补,构成现代 NLP/ML 标准 pipeline。


如需具体场景示例(如“如何用 IterableDataset 处理超大文件”或“自定义 Dataset 读取图像”),欢迎继续提问!

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐