🏠 总目录📚 本教程 08 · DataLoader 与多进程 ← →
📑 本页目录(点开跳转)

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__ 是在别的进程里执行的,不是别的线程

⚠️ 由此推出的三条实用结论:


🎲 二、随机数:先破一条过期的传言

网上流传很广的一条是「多个 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 极难发现:

✅ 两种修法,选一种:

修法 怎么写 适用
⭐ 别把发生器存进 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"]))

实跑输出:

流程图

① 三个 (2,) 张量→torch.Size([3, 2])
② python int→tensor([1, 2, 3]) torch.int64
③ python float→torch.float64
④ dict→{'x': tensor([[0., 0.],
[1., 1.]]), 'y': tensor([1, 2])}
⑤ 字符串→['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 跑一遍。

⚠️ 但要记得这个办法本身会掩盖第三节那个 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

✅ 检查点

  1. num_workers=2 时,__getitem__ 里 self.n_read += 1,主进程读到的 ds.n_read 是多少?为什么?
  2. 由「worker 拿到副本」能直接推出哪三条不该做的事?
  3. 「多个 worker 会生成完全相同的随机增强」这条传言,在 torch 2.13 上还成立吗?实测了哪三套发生器?
  4. 那个坑今天以什么形态存在?怎么用一行输出证明它?
  5. worker 的种子是怎么来的?get_worker_info() 在主进程里返回什么?
  6. 第三节那个 bug 为什么极难发现(至少说三条)?
  7. IterableDataset 配 num_workers=3 会发生什么?为什么映射式 Dataset 不会?
  8. 默认 collate_fn 对张量、int、float、dict、字符串分别做什么?哪一条会静默伤人?
  9. stack expects each tensor to be equal size 说明什么?自己写 collate 时除了补齐还必须带出什么?
  10. DataLoader 出怪事时先做什么?这个办法有什么盲区?
👀 答案
  1. 0。实测 num_workers=0 时是 8、num_workers=2 时是 0,而各条记录看到的是 [1,2,3,4,1,2,3,4] —— 两个 worker 各拿一份副本、各从 0 开始数,主进程那一份根本没被碰过。
  2. ①不要用 Dataset 属性做统计(永远是初始值,要统计就写进返回值带出来)②不要在 __getitem__ 里写缓存(每 worker 一份,内存翻 num_workers 倍且命中率低)③不要在 __init__ 里开文件句柄/数据库连接(多进程共用一个偏移量会读出乱数据),改成在 __getitem__ 里惰性打开。
  3. 不成立了。实测 torch.randint / np.random.randint / random.randint 三套全局发生器在 num_workers=2 下都是 8/8 不重复,PyTorch 逐 worker 重新播种了它们。
  4. 换成了「把发生器存进 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。
  5. DataLoader 每个 epoch 抽一个 base seed,给第 i 个 worker 发 base_seed + i —— 实测两个 worker 的 seed 是连号(…44167 / …44168),所以既互不相同、又能被 torch.manual_seed() 完整复现。get_worker_info() 在主进程里返回 None。
  6. ①数据看起来是随机的,只是周期性重复 ②训练照常收敛,只是差一点 ③num_workers=0 调试时完全正常,而这正是排查的第一反应 ④重复周期等于 num_workers,改 worker 数「随机性」就变了,容易被误判成调参有效。
  7. 整个数据流被重复 3 遍:实测 6 条的数据集在 num_workers=3 下吐出 18 条(num_workers=2 吐 12 条)。因为 IterableDataset 是数据集自己吐数据,没人告诉它该吐哪一段;映射式 Dataset 由主进程发下标,worker 之间天然不重叠。修法是在 __iter__ 里读 get_worker_info() 自己分片(i % n == wid,或更好的 files[wid::n] 在文件层面切)。
  8. 张量 stack 一层、int → int64 张量、float → float64 张量 ⚠️、dict/tuple 递归进去、字符串原样留在 list 里。会静默伤人的是 float:标签变 float64 之后整条 loss 链路被提升成双精度(实测 x 是 float32、y 是 float64、相减后 float64),不报错,GPU 上表现为「开了混合精度却没提速」。
  9. 说明样本是变长的(实测 got [3] at entry 0 and [5] at entry 1)。除了 pad_sequence 补齐,必须把真实长度一起返回,否则下游造不出 mask,padding 会被当真数据算进 loss。
  10. 先把 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

打卡记录保存在你的浏览器里,首页能看到总进度