📑 本页目录(点开跳转)
08 · DataLoader:worker 手里的不是你那个 Dataset
⏱ 72 分钟 | ⭐ 实测:num_workers=2 之后,self.n_read 在主进程里永远是 0,两个 worker 各自数到 4
🎯 一句话
num_workers > 0 的那一刻,你的 Dataset 就不再只有一份了 —— 每个 worker 进程拿到的是一份【副本】。
副本上改的东西主进程看不见,副本里的随机数发生器也是复制过去的(这一条会直接让你的数据增强失效)。
⚠️ 至于「六个参数怎么调才能喂饱 GPU」,《AI基础设施》22 · 数据管线与存储已经讲透了,这一章一个字都不重复;
「进程为什么要复制、fork 和 spawn 差在哪、Windows 上为什么必须写 if __name__」是《Python 会咬你的地方》并发那一章的正题。
这一章只讲一件事:数据集被复制之后,语义上发生了什么。
🧬 一、⭐ 副本:实测
import os
import torch
from torch.utils.data import Dataset, DataLoader
class Counting(Dataset):
def __init__(self):
self.n_read = 0 # ⚠️ 想统计「读了多少条」
def __len__(self):
return 8
def __getitem__(self, i):
self.n_read += 1
return torch.tensor([float(i), float(self.n_read), float(os.getpid())])
if __name__ == "__main__": # ⭐ Windows 上必须有这一行,原因见 Python 板块
for nw in (0, 2):
ds = Counting()
rows = [b for b in DataLoader(ds, batch_size=4, num_workers=nw)]
pids = sorted({int(v) for b in rows for v in b[:, 2].tolist()})
print(f"num_workers={nw}: 主进程里 ds.n_read = {ds.n_read}"
f" 各条记录看到的 n_read = {[int(v) for b in rows for v in b[:, 1].tolist()]}"
f" 进程数 = {len(pids)}")
实跑输出:
对照
num_workers=0: 主进程里 ds.n_read = 8 各条记录看到的 n_read = [1, 2, 3, 4, 5, 6, 7, 8] 进程数 = 1
num_workers=2: 主进程里 ds.n_read = 0 各条记录看到的 n_read = [1, 2, 3, 4, 1, 2, 3, 4] 进程数 = 2
⭐ 三个数字,每一个都值得盯一会儿:
| 观察 | 意思 |
|---|---|
主进程 ds.n_read = 0 |
worker 里改的属性主进程一律看不见。num_workers=0 时是 8 —— 同一份代码,两种行为 |
各条记录 [1,2,3,4,1,2,3,4] |
两个 worker 各从 0 开始数,谁也不知道对方存在 |
| 进程数 = 2 | __getitem__ 是在别的进程里执行的,不是别的线程 |
⚠️ 由此推出的三条实用结论:
- 不要用 Dataset 的属性做任何统计(读了多少条、哪些样本被跳过、类别计数)——
num_workers>0时它永远是初始值。要统计就把数写进返回值里带出来。 - 不要在
__getitem__里写缓存(self.cache[i] = ...)。每个 worker 一份缓存,内存翻num_workers倍,命中率还低。 - ⚠️ 不要在
__init__里打开文件句柄 / 数据库连接。句柄跟着副本走,多个进程共用一个偏移量的下场是读出乱数据。✅ 正确做法是在__getitem__里惰性打开,第一次用到才建,这样每个 worker 建自己的。
🎲 二、随机数:先破一条过期的传言
网上流传很广的一条是「多个 worker 会生成完全相同的随机增强,因为 numpy 的种子没被重设」。 这条在老版本上成立,在 torch 2.13 上已经不成立了。实测:
import random
import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader
class RngDS(Dataset):
def __len__(self):
return 8
def __getitem__(self, i):
return {"torch": torch.randint(0, 10000, (1,)).item(),
"numpy": int(np.random.randint(0, 10000)),
"random": random.randint(0, 9999)}
if __name__ == "__main__":
torch.manual_seed(0); np.random.seed(0); random.seed(0)
for nw in (0, 2):
out = {k: [] for k in ("torch", "numpy", "random")}
for b in DataLoader(RngDS(), batch_size=8, num_workers=nw):
for k in out:
out[k] += [int(v) for v in b[k]]
print(f"num_workers={nw}")
for k, v in out.items():
print(f" {k:7}: 去重后 {len(set(v))}/8")
实跑输出(num_workers=0 和 2 一致):
对照
num_workers=0
torch : 去重后 8/8
numpy : 去重后 8/8
random : 去重后 8/8
num_workers=2
torch : 去重后 8/8
numpy : 去重后 8/8
random : 去重后 8/8
⭐ 三套全局发生器(torch / np.random / random)都被 PyTorch 逐 worker 重新播种了。
种子怎么来的,也能直接看:
import os
import torch
from torch.utils.data import Dataset, DataLoader, get_worker_info
class WhoAmI(Dataset):
def __len__(self):
return 4
def __getitem__(self, i):
w = get_worker_info()
if w is None:
return torch.tensor([-1, -1, os.getpid(), 0])
return torch.tensor([w.id, w.num_workers, os.getpid(), w.seed % 100000])
if __name__ == "__main__":
for b in DataLoader(WhoAmI(), batch_size=2, num_workers=2):
for r in b.tolist():
print(f"worker id={r[0]} / 共 {r[1]} pid={r[2]} seed%100000={r[3]}")
某次实跑输出(pid 和 seed 每次都不同,要看的是它们之间的关系):
算一算
worker id=0 / 共 2 pid=40048 seed%100000=44167
worker id=0 / 共 2 pid=40048 seed%100000=44167
worker id=1 / 共 2 pid=33820 seed%100000=44168
worker id=1 / 共 2 pid=33820 seed%100000=44168
⭐ 两个 worker 的 seed 是连号的(…167 和 …168):DataLoader 每个 epoch 抽一个 base seed,
再给第 i 个 worker 发 base_seed + i。所以它们既互不相同,又能被 torch.manual_seed() 完整复现。
⚠️ get_worker_info() 在主进程里返回 None —— 上面那个 if w is None 分支不是防御性编程,num_workers=0 时它就是 None。
💀 三、但这个坑只是换了个形态
上面破掉的是「全局发生器」那一版。把发生器存进 self,它就跟着副本一起被复制了 —— 这才是今天真正会咬到你的形态,而且 np.random.default_rng() 正是 numpy 现在推荐的写法。
import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader, get_worker_info
class OwnRng(Dataset):
def __init__(self):
self.rng = np.random.default_rng(0) # ⚠️ 发生器存进了 self
def __len__(self):
return 8
def __getitem__(self, i):
return int(self.rng.integers(0, 10000))
def fix_seed(worker_id): # ⭐ 修法
info = get_worker_info()
info.dataset.rng = np.random.default_rng(torch.initial_seed() % 2**32)
if __name__ == "__main__":
for nw in (0, 2, 4):
v = [int(x) for b in DataLoader(OwnRng(), batch_size=4, num_workers=nw) for x in b]
print(f"num_workers={nw}: {v} 去重后 {len(set(v))}/8")
v = [int(x) for b in DataLoader(OwnRng(), batch_size=4, num_workers=2,
worker_init_fn=fix_seed) for x in b]
print(f"加了 worker_init_fn: {v} 去重后 {len(set(v))}/8")
实跑输出:
对照
num_workers=0: [8506, 6369, 5111, 2697, 3078, 409, 752, 165] 去重后 8/8
num_workers=2: [8506, 6369, 5111, 2697, 8506, 6369, 5111, 2697] 去重后 4/8
num_workers=4: [8506, 6369, 5111, 2697, 8506, 6369, 5111, 2697] 去重后 4/8
加了 worker_init_fn: [3232, 9238, 8653, 6793, 4674, 6441, 4372, 3919] 去重后 8/8
(⚠️ 具体数值是某次实跑 —— base seed 每次运行都变,稳定可复现的是「去重后 8/8」这个结论)
⭐ 8506, 6369, 5111, 2697 一字不差地重复了一遍 —— 两个 worker 各自从同一个状态开始摇。
换成真实场景就是:你的随机裁剪、随机翻转、随机 mask,在每 batch/num_workers 条样本上循环重复。
💀 为什么这个 bug 极难发现:
- 数据看起来是随机的(不是全零、不是全一),只是周期性重复;
- 训练照常收敛,只是收敛到的点差一点点;
num_workers=0调试时完全正常 —— 而人排查问题的第一反应正是把num_workers调成 0;- ⭐ 重复周期恰好等于
num_workers,改一下 worker 数,「随机性」就变了,很容易被误判成「调参有效果」。
✅ 两种修法,选一种:
| 修法 | 怎么写 | 适用 |
|---|---|---|
⭐ 别把发生器存进 self |
__getitem__ 里直接用 np.random.* / torch.rand 全局接口 |
绝大多数情况,因为第二节实测它们已经被逐 worker 播种了 |
worker_init_fn |
上面那个 fix_seed,用 torch.initial_seed()(已经是逐 worker 的)去重建发生器 |
你确实需要一个独立发生器(比如不想污染全局状态) |
🛑 读到这里可以停 —— 前半章讲完了(约 28 分钟)。 后半章还有:
IterableDataset在多 worker 下会把数据重复一遍 · 默认collate_fn悄悄做了什么 · 这一章不讲什么 回来的时候不用重读,直接从下一节接着看就行。
🌊 四、IterableDataset 在多 worker 下会把数据重复一遍
Dataset(映射式)由主进程分配下标,所以 worker 之间天然不重叠。
但 IterableDataset 是数据集自己吐数据的 —— 没人告诉它「你只该吐第几批」,于是每个 worker 都把整个流跑了一遍。
import torch
from torch.utils.data import IterableDataset, DataLoader, get_worker_info
class Stream(IterableDataset):
def __iter__(self):
for i in range(6):
yield i
class StreamFixed(IterableDataset):
def __iter__(self):
info = get_worker_info()
wid = 0 if info is None else info.id
n = 1 if info is None else info.num_workers
for i in range(6):
if i % n == wid: # ⭐ 自己按 worker 分片
yield i
if __name__ == "__main__":
for nw in (0, 2, 3):
v = [int(x) for b in DataLoader(Stream(), batch_size=6, num_workers=nw) for x in b]
print(f"num_workers={nw}: {v} 共 {len(v)} 条(数据集只有 6 条)")
v = [int(x) for b in DataLoader(StreamFixed(), batch_size=6, num_workers=3) for x in b]
print(f"自己分片后 num_workers=3: {sorted(v)} 共 {len(v)} 条")
实跑输出:
对照
num_workers=0: [0, 1, 2, 3, 4, 5] 共 6 条(数据集只有 6 条)
num_workers=2: [0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5] 共 12 条(数据集只有 6 条)
num_workers=3: [0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5, 0, 1, 2, 3, 4, 5] 共 18 条(数据集只有 6 条)
自己分片后 num_workers=3: [0, 1, 2, 3, 4, 5] 共 6 条
⚠️ 一个 epoch 的样本数直接乘以 num_workers,而且每条都出现了 num_workers 次。
它不报错,只是你的 epoch 变长了、每条数据被重复训练了。
⭐ 判据一句话:映射式 Dataset 是「主进程发号,worker 取货」;IterableDataset 是「worker 自己去源头取」,所以分片必须你自己写。
分片有两种写法:按下标取模(上面那种,简单但每条数据都要被读到才丢掉),或者在文件/分片层面切(files[wid::n],⭐ 更快,因为根本不去读别的 worker 那份文件)。
(流式数据怎么洗牌、缓冲区该开多大,在 《AI基础设施》22。)
📦 五、默认 collate_fn 悄悄做了什么
DataLoader 从 __getitem__ 拿回来的是一个 list 的单条样本,把它们变成一个 batch 的那一步叫 collate。默认实现的规则可以直接看:
import torch
from torch.utils.data.dataloader import default_collate
print("① 三个 (2,) 张量 →", default_collate([torch.zeros(2), torch.ones(2), torch.ones(2) * 2]).shape)
print("② python int →", default_collate([1, 2, 3]), default_collate([1, 2, 3]).dtype)
print("③ python float→", default_collate([1.0, 2.0]).dtype)
print("④ dict →", default_collate([{"x": torch.zeros(2), "y": 1},
{"x": torch.ones(2), "y": 2}]))
print("⑤ 字符串 →", default_collate(["a", "b", "c"]))
实跑输出:
流程图
四条规则:张量 stack 一层(多出一维在最前)、数字变张量、dict/tuple 递归进去、字符串原样留在 list 里不动。
⚠️⚠️ 第 ③ 条是个会静默伤人的坑:float → float64,而模型参数是 float32。
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
class FloatDS(Dataset):
def __len__(self):
return 4
def __getitem__(self, i):
return torch.randn(3), float(i) # ⚠️ 标签是 python float
if __name__ == "__main__":
x, y = next(iter(DataLoader(FloatDS(), batch_size=4)))
print("特征 dtype:", x.dtype, " 标签 dtype:", y.dtype)
print("相减之后:", (nn.Linear(3, 1)(x).squeeze() - y).dtype)
实跑输出:
要点
特征 dtype: torch.float32 标签 dtype: torch.float64
相减之后: torch.float64
⭐ 不报错,只是整条 loss 链路被静默提升成了双精度 —— CPU 上只是慢一点,GPU 上则是「明明开了混合精度却没提速」。
✅ 修法就一个字:__getitem__ 里返回 torch.tensor(float(i)),或者在 collate 之后 .float()。
两种 collate 报错(都实跑抄下来):
RuntimeError: stack expects each tensor to be equal size, but got [3] at entry 0 and [5] at entry 1
⭐ 看到 stack expects each tensor to be equal size 就是变长样本(文本、音频、点云),要自己写 collate:
import torch
from torch.nn.utils.rnn import pad_sequence
def pad_collate(batch):
lens = torch.tensor([len(s) for s in batch])
return pad_sequence(batch, batch_first=True), lens
x, lens = pad_collate([torch.ones(3), torch.ones(5) * 2])
print(x.shape, x.tolist(), lens.tolist())
实跑输出 torch.Size([2, 5]) [[1.0, 1.0, 1.0, 0.0, 0.0], [2.0, 2.0, 2.0, 2.0, 2.0]] [3, 5]。
⭐ 一定要把真实长度也带出来,否则下游没法造 mask,padding 会被当成真数据算进 loss。
TypeError: default_collate: batch must contain tensors, numpy arrays, numbers, dicts or lists; found <class '__main__.Foo'>
⭐ 这条是「你返回了个自定义对象」——PIL.Image、dataclass、spacy 的 Doc 都会撞上。
🚦 六、这一章不讲什么
| 问题 | 去哪 |
|---|---|
num_workers / pin_memory / prefetch_factor / persistent_workers / drop_last 怎么调 |
《AI基础设施》22 · 数据管线与存储 ⭐ 六个参数逐个讲透,还有「五分钟确诊数据管线是不是瓶颈」 |
进程是怎么被创建的、fork 和 spawn 差在哪、为什么 Windows 上必须写 if __name__ == "__main__"、pickle 要花多少代价 |
《Python 会咬你的地方》 并发那一章 |
| 数据该怎么存(小文件为什么是灾难、memmap、WebDataset) | 《AI基础设施》22 |
| 数据集本身干不干净、标注一不一致 | 《数据这一关》 |
⭐ 一个可以带走的排查顺序:DataLoader 出怪事时,先把 num_workers 设成 0 跑一遍。
- 设成 0 就正常 → 问题在副本语义(本章一到四节);
- 设成 0 还是不对 → 问题在
Dataset/collate_fn本身(第五节),和多进程无关。
⚠️ 但要记得这个办法本身会掩盖第三节那个 bug(num_workers=0 时随机数完全正常)—— 所以「设成 0 就好了」不等于「多进程有 bug」,也可能是你的代码依赖了只有单进程才成立的假设。
🔗 这一章连到哪里
| 相关的地方 | 为什么 |
|---|---|
07 · train() / eval() 改了什么 |
上一章末尾那条 Expected more than 1 value per channel 的诱因就是这里的 drop_last=False:最后一个 batch 只剩 1 条 |
| 09 · 自己写一个 autograd 算子 | 数据进来之后就轮到算子了;那一章讲怎么给框架加一个它没有的操作 |
| 《AI基础设施》22 · 数据管线与存储 | ⭐ 六个参数怎么调、怎么确诊数据管线是不是瓶颈、数据该用什么格式存 —— 本章刻意不碰这些 |
| 《Python 会咬你的地方》 | ⭐ 进程模型本身:fork vs spawn、为什么 Windows 必须写 if __name__、跨进程传数据的 pickle 代价。本章只讲「复制之后语义变了什么」 |
| 《数据这一关》13 · 评测集也是数据 | 第三节那个「增强周期性重复」属于数据质量问题的框架侧成因;那边讲数据本身的质量怎么量 |
| 《机器学习与深度学习基础》15 · PyTorch 实战手册 | 「GPU 利用率忽高忽低就是数据加载没配好」那一句的出处,配置方法在 AI基础设施 22 |
✅ 检查点
num_workers=2时,__getitem__里self.n_read += 1,主进程读到的ds.n_read是多少?为什么?- 由「worker 拿到副本」能直接推出哪三条不该做的事?
- 「多个 worker 会生成完全相同的随机增强」这条传言,在 torch 2.13 上还成立吗?实测了哪三套发生器?
- 那个坑今天以什么形态存在?怎么用一行输出证明它?
- worker 的种子是怎么来的?
get_worker_info()在主进程里返回什么? - 第三节那个 bug 为什么极难发现(至少说三条)?
IterableDataset配num_workers=3会发生什么?为什么映射式 Dataset 不会?- 默认
collate_fn对张量、int、float、dict、字符串分别做什么?哪一条会静默伤人? stack expects each tensor to be equal size说明什么?自己写 collate 时除了补齐还必须带出什么?DataLoader出怪事时先做什么?这个办法有什么盲区?
👀 答案
- 0。实测
num_workers=0时是 8、num_workers=2时是 0,而各条记录看到的是[1,2,3,4,1,2,3,4]—— 两个 worker 各拿一份副本、各从 0 开始数,主进程那一份根本没被碰过。 - ①不要用 Dataset 属性做统计(永远是初始值,要统计就写进返回值带出来)②不要在
__getitem__里写缓存(每 worker 一份,内存翻num_workers倍且命中率低)③不要在__init__里开文件句柄/数据库连接(多进程共用一个偏移量会读出乱数据),改成在__getitem__里惰性打开。 - 不成立了。实测
torch.randint/np.random.randint/random.randint三套全局发生器在num_workers=2下都是 8/8 不重复,PyTorch 逐 worker 重新播种了它们。 - 换成了「把发生器存进
self」——self.rng = np.random.default_rng(0)会跟着副本一起被复制。实测num_workers=2输出[8506, 6369, 5111, 2697, 8506, 6369, 5111, 2697],去重后只有 4/8;num_workers=4同样是 4/8。 - DataLoader 每个 epoch 抽一个 base seed,给第
i个 worker 发base_seed + i—— 实测两个 worker 的 seed 是连号(…44167/…44168),所以既互不相同、又能被torch.manual_seed()完整复现。get_worker_info()在主进程里返回None。 - ①数据看起来是随机的,只是周期性重复 ②训练照常收敛,只是差一点 ③
num_workers=0调试时完全正常,而这正是排查的第一反应 ④重复周期等于num_workers,改 worker 数「随机性」就变了,容易被误判成调参有效。 - 整个数据流被重复 3 遍:实测 6 条的数据集在
num_workers=3下吐出 18 条(num_workers=2吐 12 条)。因为IterableDataset是数据集自己吐数据,没人告诉它该吐哪一段;映射式 Dataset 由主进程发下标,worker 之间天然不重叠。修法是在__iter__里读get_worker_info()自己分片(i % n == wid,或更好的files[wid::n]在文件层面切)。 - 张量
stack一层、int →int64张量、float →float64张量 ⚠️、dict/tuple 递归进去、字符串原样留在 list 里。会静默伤人的是 float:标签变float64之后整条 loss 链路被提升成双精度(实测x是float32、y是float64、相减后float64),不报错,GPU 上表现为「开了混合精度却没提速」。 - 说明样本是变长的(实测
got [3] at entry 0 and [5] at entry 1)。除了pad_sequence补齐,必须把真实长度一起返回,否则下游造不出 mask,padding 会被当真数据算进 loss。 - 先把
num_workers设成 0 跑一遍:设成 0 就正常 → 问题在副本语义;设成 0 还是不对 → 问题在Dataset/collate_fn本身。⚠️ 盲区是它会掩盖第三节那个随机数 bug(单进程下随机数完全正常),所以「设成 0 就好了」也可能意味着你的代码依赖了只有单进程才成立的假设。
🛑 可以停在这里
⚡ 走神救援
⭐
num_workers > 0的那一刻,你的 Dataset 就有了 N 份副本。 实测在__getitem__里累加计数,主进程读到的是 0,而各个 worker 各数各的——__getitem__跑在别的进程里。由此三条禁令:别用 Dataset 属性做统计(要统计就写进返回值带出来)、别在
__getitem__里写缓存(内存翻 N 倍)、⚠️ 别在__init__里开文件句柄或数据库连接(多进程共用偏移量会读出乱数据,改成惰性打开)。⭐ 网上那条「多 worker 会生成相同增强」在新版本上已经不成立——PyTorch 逐 worker 播种,三套全局发生器都不重复。💀 但这个坑换了个形态:把发生器存进
self(正是 numpy 现在推荐的写法)就会跟着副本被复制,于是增强在每batch/num_workers条上循环重复。⚠️ 它难发现是因为:数据看着仍是随机的、训练照常收敛、⭐
num_workers=0调试时完全正常、改 worker 数「随机性」就变(容易误判成调参有效)。⚠️
IterableDataset在多 worker 下会把整个流重复一遍(不报错,只是 epoch 变长、每条被重复训练)——因为映射式是「主进程发号、worker 取货」,Iterable 是「worker 自己去源头取」,分片必须你自己写。⚠️ 默认 collate 有个静默陷阱:Python
float会变成float64,整条 loss 链路悄悄变成双精度。变长数据要自己 pad,⭐ 并把真实长度一起带出来,否则造不出 mask。
下一节 👉 09-自定义算子.md