
Infinibatch 实战指南基于 Kosmos-2 仓库理解大规模数据集的随机加载与可检查点迭代器【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilmInfinibatch 是一个专为深度神经网络训练设计的可检查点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 有两点本质区别它提供的是面向机器学习随机化批量数据加载的专用迭代器所有迭代器都支持检查点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.gz、corpus_chunks/corpus.01.txt.gz、corpus_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(encodingutf-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, trainTrue, seedNone, shuffleTrue, use_windowedFalse, transformNone, prefetchFalse, num_instances1, instance_rank0 ) - CheckpointableIterator各参数含义结合源码如下参数说明chunk_refs块文件的引用列表如路径名例如glob.glob(corpus_chunks/*.txt.gz)read_chunk_fnfunction(chunk_ref) - Iterator把块内容读取为条目的迭代器例如读文件并按行切分buffer_size用于洗牌的缓冲条目数默认2**20约 104 万源码注释中给出该默认值trainTrue时块按步幅strided方式分配给各实例、数据以无限排列重复False时块按连续块方式切分给各实例、数据不重复用于推理seed随机种子或Noneshuffle是否洗牌当trainFalse时必须为False否则抛ValueErrortransform对每条数据项应用的变换函数transform(Any) - Anyprefetch为True时插入一个带buffer_size的预取迭代器use_windowed临时选项切换回旧的WindowedShuffleIterator默认Falsenum_instances数据集实例数用于分布式训练中的多进程数据加载instance_rank当前实例的 rank与num_instances配合使用内部流水线结构从源码datasets.py可以看到这个便捷函数实际是把多条基础迭代器组装成一条流水线create_source_iterator(chunk_refs, train..., ...)——生成块引用序列训练时使用InfinitePermutationSourceIterator无限地生成块的排列每轮重排永不耗尽推理时使用ChunkedSourceIterator把块列表按 rank 切成连续段只服务本 rank 的那段SelectManyIterator(source_iteratorchunks, collection_selectorread_chunk_fn)——把块展平为条目对每个块调用read_chunk_fn得到条目迭代器并逐条产出若prefetchTrue插入PrefetchIterator(samples, buffer_size)若shuffleTrue默认套上BlockwiseShuffleIterator块级洗牌旧路径为BufferedShuffleIterator若提供transform追加MapIterator(samples, transform)返回流水线末端的迭代器。其中InfinitePermutationSourceIteratoriterators.py的实现值得关注它完整持有source_items这里是块路径列表只占很小内存每轮用random.shuffle生成新排列并通过记录random_state与index实现精确检查点多实例场景下它按num_instances步幅取元素保证不同 GPU/进程拿到的是互补且不重叠的数据。为保证多实例下 RNG 状态一致即使中间存在未被实际使用的排列也会完整生成见_reshuffle_as_necessary的注释说明。分桶批量读取BucketedReadaheadBatchIterator深度学习需要把多条数据组成 batch。NLP 中句子长度差异往往很大若按固定行数组 batchbatch 大小受最长序列限制短句 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(encodingutf-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 个 tokenbatch_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.pyBucketedReadaheadBatchIterator的构造参数还包括参数说明source_iterator数据源通常是无限数据源read_ahead为分组而预取的条目数key用户回调定义数据排序依据如len(line)batch_size整数或根据某个 batch 首条目估算 batch 大小的回调boundary_key可选回调将条目映射为 keykey 一旦变化就开启新 batch从而保证 batch 内所有条目的 key 相同key 不允许为Noneshuffle传False不随机化 batch 顺序默认Trueseedbatch 洗牌的随机种子其核心算法_create_batches为把预取的read_ahead条条目按key稳定降序排序稳定排序保证除了长度分组外不会破坏此前的随机化然后依次聚合成 batch每个 batch 满batch_size函数版则按首条目动态计算即收尾若提供了boundary_keykey 变化时强制开启新 batch。每轮预取窗口内形成的 batch 列表还会被随机打乱后再产出。检查点由source_state预取窗口起始处的数据源状态、random_state与num_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)这里用到的MapIteratoriterators.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_processesnum_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_tensor用utils.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 --failfasttest_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),仅供参考