Skip to content

深度学习的分布式训练与大数据集加载

2026-08-16 · 2930 字 · 10 分钟 · 浏览量

引入(数据集<内存<硬盘)

刚接触深度学习的时候,我们的设备是单机单卡,且总是默认数据集可以完整加载到内存中。彼时采样总是在全局进行,也就是说每个 epoch 中,任意两个样本被采样的概率是相等的并且都可能进入同一批次。

此时我们使用的是 torch.utils.data.Datasettorch.utils.data.DataLoaderDataset 是可以随机访问的,即任意指定 idx,都可以读取 dataset[idx] 对应的数据条目。例如:

python
class InMemoryDataset(torch.utils.data.Dataset):
    def __init__(self, dataset_name):
        with open(dataset_name, 'rb') as f:
            self.dataset = f.read()  # 假设这里已经读成了一个连续的数组
    def __getitem__(self, idx):
        return self.dataset[idx]

数据集采样通过 DataLoadershuffle=True 实现,它本质上是使用了 RandomSampler 来随机产生索引序列,然后按照 batch_size 切分成小段,并使用 Dataset.__getitem__() 来读取小段中的索引。例如 RandomSampler 产生 [1,3,4,2,6,8,0,5]batch_size=4,那么第一个 batch 则读取 dataset[1],dataset[3],dataset[4],dataset[2] 并合并为一个 batch。

内存映射 memory-map 加速读取(内存<数据集<硬盘)

有时数据并没有在数据集创建时就加载到内存中,或者说无法把全体数据集都加载到内存中,此时 Dataset.__getitem__() 可能写成以下形式:

python
class FileDataset(torch.utils.data.Dataset):
    def __init__(self, file_list):
        self.file_list = file_list
    def __getitem__(self, idx):
        file_name = self.file_list[idx]
        with open(file_name, 'rb') as f:
            data = f.read()
        return data

也就是说,只有在数据集具体索引时,数据才从文件系统中读入内存。但是,f.read() 内核会做这么几件事:

  1. 从硬盘将文件读到内核缓冲区
  2. 从内核缓冲区拷贝到用户缓冲区

注意到,第二步有一次数据拷贝的开销。这主要因为一般来说内核缓冲区和用户缓冲区是隔离的,内核缓冲区可能随内核任务调度而改变内容,而用户缓冲区则是单独为用户开辟的内存空间,不会被非用户进程写入。

而内存映射则省去了第二步的拷贝,它直接把内核缓冲区的地址告诉用户,从而让用户直接前往相应的缓冲区获取数据,当用户申请读取这一块缓冲区时,内核会检查上面是否有用户申请的内容,如果没有则从硬盘中读取。由于内存映射只是映射地址,而没有实际占用一块内存空间,因此完全可以为体积大于物理内存空间的数据集设置内存映射,只要保证每次读取时所需要的那一小批数据量小于内存空间即可,即有足够的缓冲区来存放一次读取的数据量。加入缓冲区的这部分数据在内存中的缓存时间也是随内核调度的,如果内核资源紧张,这部分缓冲区中的数据会换出,需要时重新加载。因此如果当物理内存总量很小时,缓存会频繁换出,因而会频繁从硬盘读取文件,所以说第一步的 IO 开销是无法节省的。

Arrow 格式就是一种基于 mmap 的数据集格式。我们可以使用 HF 的 datasets 库来使用这种格式,它提供了 torch.utils.data.Dataset 这一抽象接口更具体的实现。

多文件还是单文件

当数据量上涨后,另一项过程的开销也变得显著,这就是路径解析/寻址,即 with open(file_name) as f 的开销,这一部分需要触碰文件系统的 inode 树。除此以外,文件系统本身存储文件时,也会存储除了文件内容以外的元数据。因此把一系列数据条目保存成许多小文件和单个大文件就会产生区别。

像 Arrow 这样的大文件格式,只需要在最初 open()(并 mmap)一次,之后访问其中任意一条数据记录都不再需要触碰文件系统,而是通过 Footer 记录的 offset 直接在已经映射的内存地址上做算术定位。这个查找过程是一次纯内存操作,不涉及磁盘/文件系统层面的路径解析,因此文件数量再多(即样本条目再多),也不会像"每个样本一个独立小文件"那样,随着规模增长而线性增加 open() 调用次数和路径解析开销。

随机访问与顺序访问

这里我们要区分两种数据访问方式,随机访问与顺序访问。前者意味着给定索引可以以 O(1) 的时间复杂度访问数据,而后者只能以顺序方式遍历访问,这也就意味着给定索引需要 O(n) 的时间复杂度访问数据。

在 torch 中这两种数据访问方式对应着两种数据集接口格式 DatasetIterableDataset,二者分别带有 __getitem__() 方法和 __iter__() 方法。

注意,随机访问和可枚举是两码事,例如 tar -t 可以枚举出其中包含的文件名列表,但是这并不意味着它访问某个文件的时间是 O(1)。随机访问的含义是可以在 O(1) 的时间内找到内容,而顺序访问不能,它必须按顺序扫描直至遇到目标文件。

网络数据集 WebDataset 加载超大数据集(内存<硬盘<数据集)

当数据量继续增长,超过本地存储的大小时,我们就只能以网络通信的方式访问云存储了。云存储与本地存储不同,它不支持随机访问,只支持顺序访问。顺序访问的数据格式通常是.tar,我们可以使用webdataset库来使用这种格式。

本地磁盘的随机访问 seek() 由文件系统原生支持、开销很小(微秒级),而云存储理论上可以通过 HTTP Range 请求实现随机访问,但每一次 Range 请求都是一次独立的网络往返,延迟通常是本地磁盘随机访问的百倍到千倍,而且大多数压缩格式(如 .tar.gz)本身也无法从中间任意位置直接解压。因此实践中云存储被当作只适合顺序访问的介质。

采样

随机访问还是顺序访问对训练最大的影响在于采样策略。随机访问可以给定索引访问数据,因此对采样策略没有任何限制,任何返回特定索引的策略都是可行的。但是对于顺序访问的数据集而言,采样策略就没有办法任意地给出索引了。

为了尽可能保证数据集采样的随机性,避免引入任何顺序造成的偏差,同时又兼顾顺序访问的效率,人们把一个大数据集拆分成若干分片(shard),因而对不同 shard 的访问是随机的(相当于不同文件),而同一个 shard 内部仍然保持顺序访问。此外,即便每个 shard 是顺序访问的,我们在顺序读取一批数据到 buffer 中后,也可以在 buffer 中对数据进行随机访问(因为此时已经进入内存了)。

分布式训练

以上我们讲的还都是单机单卡训练,现在我们讲分布式训练。分布式训练,就是在多台 GPU 上同时训练同一个模型。这里我们只讨论最简单的数据并行的分布式训练。对于数据并行的分布式训练而言,每台 GPU 都加载一份完整的模型副本与不同的数据子集,通过梯度平均来更新模型参数。如果忽略梯度平均的通信开销,数据并行的分布式训练可以近似认为是线性加速的,也就是说如果有 N 台 GPU,那么训练速度大约是单台 GPU 的 N 倍。

对于单机多卡而言,通信是通过 NVLink 进行的,通信开销可以忽略不计。单机多卡与多机多卡的区别在于主要通信方式不同,单机多卡是通过 NVLink 进行通信,而多机多卡是通过网络进行通信,通信开销较大。由于我们组多机多卡的通信一直有些问题,我很难调试成功,这里不讨论。

数据并行

单机多卡的分布式训练在 pytorch 通过 torch.nn.parallel.DistributedDataParallel(DDP)实现。具体来说,如果用 N 张 GPU 分布式训练,DDP 会创建 N 个进程,每个进程绑定一张 GPU、若干 CPU 和独立的内存空间,并且每个进程加载一份完整的模型副本,然后每个进程加载各自的数据。

分布式的数据采样由 torch.utils.data.distributed.DistributedSampler 实现。DistributedSampler 会根据进程数和进程编号来划分数据集的子集,并且以 epoch 数为随机种子打乱数据集。例如,我们使用 2 张 GPU 进行分布式训练,DistributedSampler 产生 [1,3,4,2,6,8,0,5],它将 [1,4,6,0] 交给 GPU0 所在的进程,而将 [3,2,8,5] 交给 GPU1 所在的进程。

注意,如果在 Dataset 的 __init__() 阶段加载数据集,很可能整个数据集加载 N 次。因为相同的代码被每个进程都执行了一次。即便采样的索引被分给了不同的进程,每个进程是独立地加载自己内存空间中的数据集,只不过是取了其中被分配的索引位置。为了避免内存浪费,我们可以把数据加载放到 __getitem__() 再进行。

又或者我们把所有的数据都放在一个进程共享的 buffer 中,而不是单独的进程空间。mmap 数据集就可以实现这一点(本质上就是 __getitem__ 时的加载)。

再或者,在 __init__() 阶段就安排好根据进程的 rank 数加载相应的数据分片,例如 GPU0 加载所有偶数次分片,GPU1 加载奇数次分片,这样就避免了内存重复加载数据集的浪费。但这就使得我们不可能用一个全局的 DistributedSampler 来采样,因为不同分片的索引是不共享的。此时可以使用 WebDataset 的采样方式。它提供了一套组合式的 pipeline 方法,用来在没有全局 DistributedSampler 的情况下尽量还原随机性:

py
dataset = (
    wds.WebDataset(shard_urls, shardshuffle=True)  # 分片顺序打乱
    .shuffle(1000)                                  # 分片内部用大小为 1000 的 buffer 局部打乱
)
  • shardshuffle=True:每个 epoch 开始前打乱分片的顺序,替代"样本级别全局打乱",做到"分片级别打乱"。
  • .shuffle(buffer_size):分片内部仍是顺序读取,但读进一个固定大小的 buffer 后再随机吐出,buffer 越大越接近真随机,也越占内存。
  • 多 GPU/多 worker 场景下,WebDataset 内部会自动按 rank 和 worker 对分片做切分(nodesplitter / splitter 机制),确保每个分片只被一个 worker 消费一次,不需要手动写 DistributedSampler
  • 为了避免各 rank 因分片数据量不均导致 DDP 卡死,通常会配合 .with_epoch(n) 固定每个 epoch 的 step 数,数据不够就循环回绕补齐。

本质上还是"分片顺序打乱 + 分片内 buffer 局部打乱"这两层组合,只是 webdataset 库把这套逻辑封装成了现成的 pipeline,不需要自己实现。

返回

人同此心,心同此理;如风沐面,若水润心