Infinibatch 实战指南:基于 Kosmos-2 仓库理解大规模数据集的随机加载与可检查点迭代器
2026/9/14 4:35:41 网站建设 项目流程

Infinibatch 实战指南:基于 Kosmos-2 仓库理解大规模数据集的随机加载与可检查点迭代器

【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm

Infinibatch 是一个专为深度神经网络训练设计的可检查点(checkpointable)迭代器库,用于对远超内存容量的海量数据集进行随机化数据加载。本文以 Kosmos-2 仓库中的 Infinibatch 子目录 为对象,完整讲解从数据分块、随机读取、分桶批量到接入深度学习框架的实战流程,并结合仓库源码剖析其分层洗牌、100% 精确断点续训、多 GPU 数据切分与预取等底层原理,帮助读者掌握在大规模语料上构建高效数据管线的完整能力。

Infinibatch 是什么

Infinibatch 是一组可检查点的 Python 迭代器集合,专门用于深度神经网络训练中大规模数据集的随机化加载。与常见的DataLoader思路不同,它把"读取数据"抽象为一条由多个迭代器组成的流水线(pipeline),每个迭代器只做一件事,且全部支持断点保存与恢复。

其核心特性(见 README.md 与 包内__init__.py)包括:

  • 支持远超 RAM 容量的语料库;
  • 采用分层(hierarchical)块级 + 句子级两级随机化,覆盖整个语料,且每个 epoch 的随机化结果不同;
  • 只加载当前需要的数据;
  • 启动极快(无需预读完整语料);
  • 数据准备成本极低(无需构建索引);
  • 多 GPU 场景下,每个 GPU 只加载自己所需的数据;
  • 100% 精确的检查点,恢复时无需重读检查点之前的所有数据;
  • 支持自动分桶(bucketed batching)与动态 batch size;
  • 内置预取线程/进程;
  • 迭代器可组合,支持负采样等多文档复杂批处理场景。

安装与环境要求

Infinibatch 要求 Python 3.6 或更高版本,且没有任何第三方依赖(README 明确指出 "has no dependencies")。截至当前仓库版本,还没有对应的 pip 包发布。

仓库内 setup.py 定义的包名为infinibatch,版本号为0.1.0,仅包含find_packages()发现的全部子包,不含任何install_requires,这印证了"零依赖"的声明。本地安装方式:

git clone <仓库地址> cd <仓库根目录>/kosmos-2/infinibatch pip install -e .

其中-e表示以可编辑(开发)模式安装,便于直接修改源码后立即生效。

核心概念一:迭代器与惰性求值

Infinibatch 以 Python 标准迭代器协议为基础:一个迭代器表示一条数据流,可通过for循环或反复调用next()逐条取出数据。

迭代器对数据类型完全无感——数据项的具体类型由用户提供的读取函数决定:NLP 场景中通常是文本元组,其他场景可以是图片、带文本标注的音频文件等。这种"数据格式由用户定义"的设计,使 Infinibatch 可以服务文本、多模态等各类任务。

与 Python 标准库itertools相比(见 iterators.py 模块文档字符串),Infinibatch 有两点本质区别:

  1. 它提供的是面向机器学习随机化批量数据加载的专用迭代器;
  2. 所有迭代器都支持检查点(checkpointing),因此与itertools并不直接兼容。

由于迭代器按需惰性求值,Infinibatch 只对"正在消费的那一条数据"执行操作,而不是一次性处理整个数据集。这正是其低启动时间、低内存开销的根源。

核心概念二:检查点与断点续训

长期训练难免崩溃。Infinibatch 的迭代器是可检查点的:任何时刻都可以通过getstate()取回数据流中的当前位置(即"检查点"),之后用setstate()"回卷"到该位置。训练中每次保存中间模型时,把迭代器检查点一并落盘;崩溃后恢复时,将迭代器重置到保存的检查点,数据读取器就会产出与未崩溃时完全相同的数据项序列。

从源码看(iterators.py),这一机制由抽象基类CheckpointableIterator统一定义:

  • getstate() -> Dict:返回表示当前状态的检查点对象;在迭代器流水线中,它会递归调用上游迭代器的getstate(),因此只需在流水线最后一个迭代器上调用,即可捕获整条流水线的状态;
  • setstate(checkpoint):将流水线重置到检查点状态;传入None则重置到构造后的初始状态;
  • close():递归遍历整条流水线并关闭所有PrefetchIterator(见下文),避免悬挂进程/线程。

此外,CheckpointableIterator还实现了 Python pickle 协议(__getstate__/__setstate__),使检查点可以借助pickle模块直接序列化落盘,也可随模型检查点一起保存。测试代码 test_iterators.py 中的TestFiniteIteratorCheckpointingMixin专门验证了"取检查点→消费数据→恢复检查点→重新输出一致"这一行为。

数据准备:把语料切分为小块

使用 Infinibatch 的唯一数据组织要求是:把数据拆成大量小块(chunks)。块是从磁盘载入 RAM 的最小数据单元,Infinibatch 在内存中持有块的随机子集并从中随机取样。

一个最简单的切分方式是利用 Linux 的split命令。下面以 6 行文本、每行一条数据为例,先创建语料文件:

echo \ 'Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat. Duis aute irure dolor in reprehenderit in voluptate velit esse cillum dolore eu fugiat nulla pariatur. The quick brown fox jumps over the lazy dog.' \ > corpus.txt

再把它切成 3 个每块 2 行的 gzip 压缩块,存到新目录corpus_chunks中:

mkdir corpus_chunks split --lines 2 --numeric-suffixes \ --filter 'gzip > corpus_chunks/$FILE.txt.gz' \ corpus.txt corpus.

执行后会生成三个文件:corpus_chunks/corpus.00.txt.gzcorpus_chunks/corpus.01.txt.gzcorpus_chunks/corpus.02.txt.gz。可用以下命令校验切分结果:

zcat corpus_chunks/corpus.*.txt.gz

提示:对超大规模语料,建议用pigzapt-get install pigz)替代gzip,其多线程实现可显著提升压缩/解压速度。

随机读取:chunked_dataset_iterator()及其源码剖析

读取数据的最简单方式是利用便捷函数chunked_dataset_iterator()(位于 datasets.py)。下面这个程序随机顺序地逐条输出语料内容:

import gzip, glob from infinibatch import datasets as ds ds = ds.chunked_dataset_iterator( chunk_refs = glob.glob('corpus_chunks/corpus.*.txt.gz'), read_chunk_fn = lambda path: iter(gzip.decompress(open(path, "rb") \ .read()).decode(encoding='utf-8') \ .splitlines()), buffer_size = 6, seed = 1) for i in range(10): print(next(ds))

输出为 6 个例句的随机排列(示例输出):

Lorem ipsum dolor sit amet, consectetur adipiscing elit, Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat. Duis aute irure dolor in reprehenderit in voluptate velit esse cillum dolore eu fugiat nulla pariatur. The quick brown fox jumps over the lazy dog. sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. consectetur adipiscing elit, Lorem ipsum dolor sit amet, The quick brown fox jumps over the lazy dog. sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.

注意:buffer_size决定任一时刻读入内存用于随机取样的句子数。在数亿行文本的真实场景中,该参数应设置为数百万量级;内存占用与启动时间与 buffer 大小成正比(但仍远低于把整个语料载入内存)。

函数签名与参数语义

chunked_dataset_iterator的完整签名(datasets.py)为:

def chunked_dataset_iterator( chunk_refs, read_chunk_fn, buffer_size, train=True, seed=None, shuffle=True, use_windowed=False, transform=None, prefetch=False, num_instances=1, instance_rank=0 ) -> CheckpointableIterator

各参数含义(结合源码)如下:

参数说明
chunk_refs块文件的引用列表(如路径名),例如glob.glob('corpus_chunks/*.txt.gz')
read_chunk_fnfunction(chunk_ref) -> Iterator,把块内容读取为条目的迭代器,例如读文件并按行切分
buffer_size用于洗牌的缓冲条目数,默认2**20(约 104 万),源码注释中给出该默认值
trainTrue时块按步幅(strided)方式分配给各实例、数据以无限排列重复;False时块按连续块方式切分给各实例、数据不重复(用于推理)
seed随机种子(或None
shuffle是否洗牌;当train=False时必须为False,否则抛ValueError
transform对每条数据项应用的变换函数transform(Any) -> Any
prefetchTrue时插入一个带buffer_size的预取迭代器
use_windowed临时选项,切换回旧的WindowedShuffleIterator(默认False
num_instances数据集实例数,用于分布式训练中的多进程数据加载
instance_rank当前实例的 rank,与num_instances配合使用

内部流水线结构

从源码(datasets.py)可以看到,这个便捷函数实际是把多条基础迭代器组装成一条流水线:

  1. create_source_iterator(chunk_refs, train=..., ...)——生成块引用序列:训练时使用InfinitePermutationSourceIterator(无限地生成块的排列,每轮重排,永不耗尽);推理时使用ChunkedSourceIterator(把块列表按 rank 切成连续段,只服务本 rank 的那段);
  2. SelectManyIterator(source_iterator=chunks, collection_selector=read_chunk_fn)——把"块"展平为"条目":对每个块调用read_chunk_fn得到条目迭代器并逐条产出;
  3. prefetch=True,插入PrefetchIterator(samples, buffer_size)
  4. shuffle=True,默认套上BlockwiseShuffleIterator(块级洗牌),旧路径为BufferedShuffleIterator
  5. 若提供transform,追加MapIterator(samples, transform)
  6. 返回流水线末端的迭代器。

其中InfinitePermutationSourceIterator(iterators.py)的实现值得关注:它完整持有source_items(这里是块路径列表,只占很小内存),每轮用random.shuffle生成新排列,并通过记录random_stateindex实现精确检查点;多实例场景下,它按num_instances步幅取元素,保证不同 GPU/进程拿到的是互补且不重叠的数据。为保证多实例下 RNG 状态一致,即使中间存在未被实际使用的排列,也会完整生成(见_reshuffle_as_necessary的注释说明)。

分桶批量读取:BucketedReadaheadBatchIterator

深度学习需要把多条数据组成 batch。NLP 中句子长度差异往往很大,若按固定行数组 batch,batch 大小受最长序列限制,短句 batch 会浪费 GPU 显存与算力。

BucketedReadaheadBatchIterator实现了一种**分桶(bucketing)**算法(模型参考自 Marian NMT 工具包):预读大量随机化条目(真实场景通常达数百万条,本例为 6 条),按长度排序并聚成长度相近的 batch,再以随机顺序逐个产出。

import gzip, glob from infinibatch import datasets as ds from infinibatch import iterators as it ds = ds.chunked_dataset_iterator( chunk_refs = glob.glob('corpus_chunks/corpus.*.txt.gz'), read_chunk_fn = lambda path: iter(gzip.decompress(open(path, "rb") \ .read()).decode(encoding='utf-8') \ .splitlines()), buffer_size = 6, seed = 1) bs = it.BucketedReadaheadBatchIterator( source_iterator = ds, # note: this is the iterator from above read_ahead = 6, key = lambda line: len(line), batch_size = 2, seed = 1) for i in range(25): print(next(bs))

注意BucketedReadaheadBatchIterator接受上一步的随机句子序列迭代器(ds)作为数据源——这就是 Infinibatch 的迭代器流水线组合方式(与 Pythonitertools的组合思想一脉相承)。一旦某个迭代器被传给另一个迭代器作为数据源,它即归后者所有,调用方代码不得再访问它

预期输出为 2 条一组、长度相近的随机组合:

['sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.', 'The quick brown fox jumps over the lazy dog.'] ['consectetur adipiscing elit,', 'Lorem ipsum dolor sit amet,'] ['Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat.', 'Duis aute irure dolor in reprehenderit in voluptate velit esse cillum dolore eu fugiat nulla pariatur.']

本例中"分桶组合方式没有变化"只是示例规模过小造成的假象,真实数据远大于 batch size 时不会如此。

动态 batch size:以函数作为batch_size

固定行数 batch 会浪费 GPU 资源,理想做法是:batch 中装入的行数应刚好占满 GPU 显存,即由 batch 内最长行的 token 数决定。Infinibatch 允许把batch_size传为函数:该函数接收 batch 中最长条目,估算最多能装下多少条。

以下代码假设 batch 最多容纳 150 个 token:

batch_size = lambda longest_line: 150 // len(longest_line),

输出中短句被分组、长句独立成 batch(因为组合后会超过 150 字符上限):

['consectetur adipiscing elit,', 'Lorem ipsum dolor sit amet,'] ['Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat.'] ['sed do eiusmod tempor incididunt ut labore et dolore magna aliqua.', 'The quick brown fox jumps over the lazy dog.'] ['Duis aute irure dolor in reprehenderit in voluptate velit esse cillum dolore eu fugiat nulla pariatur.']

源码级原理

从实现看(iterators.py),BucketedReadaheadBatchIterator的构造参数还包括:

参数说明
source_iterator数据源,通常是无限数据源
read_ahead为分组而预取的条目数
key用户回调,定义数据排序依据(如len(line)
batch_size整数,或根据某个 batch 首条目估算 batch 大小的回调
boundary_key可选回调,将条目映射为 key;key 一旦变化就开启新 batch,从而保证 batch 内所有条目的 key 相同(key 不允许为None
shuffleFalse不随机化 batch 顺序(默认True
seedbatch 洗牌的随机种子

其核心算法(_create_batches)为:把预取的read_ahead条条目按key稳定降序排序(稳定排序保证除了长度分组外,不会破坏此前的随机化),然后依次聚合成 batch,每个 batch 满batch_size(函数版则按首条目动态计算)即收尾;若提供了boundary_key,key 变化时强制开启新 batch。每轮预取窗口内形成的 batch 列表还会被随机打乱后再产出。检查点由source_state(预取窗口起始处的数据源状态)、random_statenum_served(当前窗口已产出的 batch 数)三者构成,恢复时精确回到当前窗口的起始处。

把 batch 转为 numpy 数组:MapIterator与自定义 collate

最后一步是把文本 batch 交给深度学习框架。典型做法是:文本 token 化后用词汇表索引表示每个 token,再 padding 到等长并转为numpy数组。下面示例以"每个字符即一个 token、ASCII 码即索引",不足部分用-1填充:

import numpy as np def collate(lines_batch): # tokenize all lines in the batch and map to unit ids ids_batch = [[ord(c) for c in line] for line in lines_batch] # create a padded numpy array as wide as the longest line, # where shorter sequences are padded with -1 width = max(len(ids) for ids in ids_batch) return np.array([ids + [-1] * (width-len(ids)) for ids in ids_batch]) bs = it.MapIterator( source_iterator = bs, transform = collate)

这里用到的MapIterator(iterators.py)会对每条数据项应用用户提供的函数或 lambda。输出为填充后的 numpy 数组,多句 batch 中较短的句子以-1补足:

[[ 99 111 110 115 101 99 116 101 116 117 114 32 97 100 105 112 105 115 99 105 110 103 32 101 108 105 116 44] [ 76 111 114 101 109 32 105 112 115 117 109 32 100 111 108 111 114 32 115 105 116 32 97 109 101 116 44 -1]] [[ 85 116 32 101 110 105 109 32 97 100 32 109 105 110 105 109 32 118 101 110 105 97 109 44 32 113 117 105 115 32 110 111 115 116 114 117 100 32 101 120 101 114 99 105 116 97 116 105 111 110 32 117 108 108 97 109 99 111 32 108 97 98 111 114 105 115 32 110 105 115 105 32 117 116 32 97 108 105 113 117 105 112 32 101 120 32 101 97 32 99 111 109 109 111 100 111 32 99 111 110 115 101 113 117 97 116 46]] [[115 101 100 32 100 111 32 101 105 117 115 109 111 100 32 116 101 109 112 111 114 32 105 110 99 105 100 105 100 117 110 116 32 117 116 32 108 97 98 111 114 101 32 101 116 32 100 111 108 111 114 101 32 109 97 103 110 97 32 97 108 105 113 117 97 46] [ 84 104 101 32 113 117 105 99 107 32 98 114 111 119 110 32 102 111 120 32 106 117 109 112 115 32 111 118 101 114 32 116 104 101 32 108 97 122 121 32 100 111 103 46 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1 -1]] [[ 68 117 105 115 32 97 117 116 101 32 105 114 117 114 101 32 100 111 108 111 114 32 105 110 32 114 101 112 114 101 104 101 110 100 101 114 105 116 32 105 110 32 118 111 108 117 112 116 97 116 101 32 118 101 108 105 116 32 101 115 115 101 32 99 105 108 108 117 109 32 100 111 108 111 114 101 32 101 117 32 102 117 103 105 97 116 32 110 117 108 108 97 32 112 97 114 105 97 116 117 114 46]]

进阶:用基础迭代器组合复杂流水线

chunked_dataset_iterator()覆盖了最常见的场景,但多任务学习等真实需求往往需要更复杂的组合,此时要用 iterators.py 中的底层构建块自行组装。该模块中的迭代器按角色可分为几类:

数据源迭代器(置于流水线最前端)

  • InfinitePermutationSourceIterator:接受列表并无限地生成其排列,训练场景的数据源首选,支持多 GPU 切分;
  • ChunkedSourceIterator:按 rank 把列表切成连续段并逐条产出,用于推理/验证场景,支持多 GPU 推理切分;
  • NativeCheckpointableIterator:把普通 Python iterable 包装成可检查点迭代器,主要用于演示与调试(恢复检查点时需逐条重放,效率较低,且不能接收迭代器)。

变换与映射

  • MapIterator:对每条数据应用变换函数;
  • ParallelMapIterator:用多进程并行执行变换(num_processes+num_items_per_process),要求变换函数可 pickle(应定义为顶层函数);
  • RecurrentIterator:以有状态 step 函数迭代(step_function(state, item) -> (new_state, output));
  • SamplingRandomMapIterator:在变换的同时传入可检查点的随机数生成器。

批量与窗口

  • FixedBatchIterator:把 N 条连续条目合成一个 batch 列表;
  • WindowedIterator:产出宽度为width的滑动窗口元组;
  • SelectManyIterator:把每条条目投影为一个序列并展平(类似 LINQ 的 SelectMany)。

组合与预取

  • ZipIterator:类似zip(),按条目逐条对齐多个迭代器(到最短者耗尽为止);
  • MultiplexIterator:用一个控制迭代器产出的索引序列,从多个输入迭代器中挑选下一个条目,是实现多数据源混合(如多任务采样)的关键;
  • PrefetchIterator:在独立进程中预取数据到缓冲队列,以隐藏上游 I/O 延迟。

值得留意的是PrefetchIterator的实现细节:它利用 UNIXfork创建预取进程,因此不支持 Windows(源码中在非 fork 系统上会退化为直接返回源迭代器并打印警告);实验版_ForkPrefetchIteratorExperimental在 iterators.py 中详细解释了为何把进程间队列容量限制为 1、把真正的缓冲放在主进程的线程安全本地队列中——这是为了避免 CPython GIL 下预取进程内多个线程(队列喂给线程、PyTorch 张量共享内存的额外线程)互相争抢导致的严重卡顿。同时,官方强调:包含PrefetchIterator的流水线必须手动调用close()来回收进程/线程资源,不能依赖垃圾回收器(CPython 不保证__del__被调用)。

在 Kosmos-2 中的实际应用

Infinibatch 并非孤立工具,它正是 Kosmos-2 等大型多模态模型训练时实际使用的数据加载基础。以 basic_loader.py 为例,其中的BaseBatchGen类直接继承infinibatch.iterators.CheckpointableIterator

  • 通过_build_iter()构建 Infinibatch 迭代器并保存在self._iter
  • getstate/setstate分别暴露为state_dict/load_state_dict,从而把数据读取检查点无缝接入 fairseq 的模型检查点体系,实现"保存模型即保存数据读取进度";
  • __next__直接透传next(self._iter)close()透传上游的close()
  • _move_to_tensorutils.apply_to_sample把 numpy batch 递归转换为torch.tensor

同一目录下的 lm_loader.py、mlm_loader.py、spm_lm_loader.py 等均以BaseBatchGen为基类,用chunked_dataset_iterator组装各自的训练数据管线;utils.py 中的ConcatIterator则用于拼接多个数据源。这说明:凡是用 Infinibatch 构建的迭代器,天然获得随机化、可检查点、多 GPU 切分三大能力,且能零成本地挂接到主流训练框架的检查点机制上

测试与文档生成

仓库为 Infinibatch 提供了完整的单元测试(test 目录),运行方式:

python -m unittest discover -s test

若希望首个失败即停止:

python -m unittest discover -s test --failfast

test_iterators.py 覆盖了各迭代器的基本功能与检查点行为(重置到起点、从起点取检查点、从任意位置取检查点后恢复并保持输出一致),并验证了多实例切分(world_sizes覆盖 1~73 种规模)的正确性;test_datasets.py 与 test_doctests.py 分别覆盖便捷数据集函数与文档示例。若安装了mypy,还可做类型检查:

mypy infinibatch

文档方面,仓库内提供了 docs/config.mako 模板;安装pdoc3后可本地预览 API 文档(pdoc --template-dir docs --http : infinibatch),合并代码前可用pdoc -o docs --template-dir docs --html infinibatch重新生成 HTML 文档。其中iterators.py模块的文档字符串本身就是进阶用法的权威说明(含完整的迭代器流水线演示与检查点往返示例),是继本文教程之后继续深入的最佳入口。

总结

Infinibatch 用一套简洁而严密的迭代器模型,优雅地解决了大模型训练中"超大数据集随机化 + 精确断点续训"这一对看似矛盾的需求:分块存储降低内存压力,分层洗牌在控制内存的同时保证随机质量,递归式检查点让断点恢复精确到单条数据,多实例步幅切分让多 GPU 各取所需,而BucketedReadaheadBatchIterator的动态分桶批量则充分榨取 GPU 算力。从 README 教程 到 datasets.py 与 iterators.py 的实现,再到 basic_loader.py 的工程集成,本仓库为读者提供了从概念到生产落地的完整闭环,可直接复用于各类大规模、多模态、多 GPU 训练任务的数据管线构建。

【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询