ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

NeMo ConcatDataset 与 ConcatMapDataset 源码级指南:多数据集温度采样、随机采样与轮询混合

2026/9/13 23:29:25 拓冰建站 浏览量
NeMo ConcatDataset 与 ConcatMapDataset 源码级指南:多数据集温度采样、随机采样与轮询混合 NeMo ConcatDataset 与 ConcatMapDataset 源码级指南多数据集温度采样、随机采样与轮询混合【免费下载链接】SpeechA scalable generative AI framework built for researchers and developers working on Large Language Models, Multimodal, and Speech AI (Automatic Speech Recognition and Text-to-Speech)项目地址: https://gitcode.com/GitHub_Trending/nem/Speech本文围绕 NeMo 框架中nemo.collections.common.data.dataset模块的核心数据类ConcatDataset与ConcatMapDataset展开深入讲解如何将多个训练数据集拼接为单一数据集并通过温度采样temperature、随机采样random、轮询round-robin三种策略控制各子数据集的采样权重。读完本文你将掌握这两个类的全部构造参数、底层迭代与分布式分片原理并能在 ASR/语音分类等真实训练配置中正确接线使用。一、为什么需要数据集拼接在语音识别、语音分类、自监督预训练等大规模训练场景中单一数据源往往无法满足需求需要将多个领域、多个说话风格、多个语言的 manifest 混合训练以提升模型泛化能力各子数据集规模差异悬殊如一个 1000 小时的英文库配一个 50 小时的小语种库直接拼接会导致小数据集被淹没需要在训练中控制不同数据源的采样比例甚至动态上下采样。NeMo 在 nemo/collections/common/data/dataset.py 中提供了三个相关类ConcatDataset、ConcatMapDataset与CodeSwitchedDataset三者均在模块的__all__导出列表中见 dataset.py#L28。其中前两者正是本文主角——它们把多个子数据集包装成一个统一数据集并按照指定策略决定每一步从哪个子数据集取样本。而文档对应的 API 页位于 docs/source/common/data.rst。二、ConcatDataset可迭代式多数据集混合采样ConcatDataset继承自torch.utils.data.IterableDataset见 dataset.py#L31这意味着它通过__iter__惰性产出样本而不是通过__getitem__按索引访问天然适合流式读取 tarred 数据集等场景。2.1 构造函数参数详解其完整签名如下见 dataset.py#L51-L62ConcatDataset( datasets: List[Any], shuffle: bool True, sampling_technique: str temperature, sampling_temperature: int 5, sampling_scale: int 1, sampling_probabilities: List[float] None, seed: Optional[int] None, global_rank: int 0, world_size: int 1, )各参数含义与行为如下表参数类型默认值说明datasetslist必填待混合采样的子数据集列表所有子数据集必须是同一类型要么全部是可迭代数据集要么全部是 map 风格数据集否则构造时抛出ValueErrorshuffleboolTrue是否打乱各子数据集内部顺序。仅对 map 风格数据集生效可迭代数据集直接使用其自身迭代器sampling_techniquestrtemperature选择从哪个子数据集取样本的策略支持temperature、random、round-robin非法值抛出ValueErrorsampling_temperatureint5仅temperature策略使用控制采样分布偏向程度sampling_scaleint1对拼接后数据集整体进行上/下采样数据集总长度乘以该系数见 dataset.py#L105-L107sampling_probabilitieslistNone仅random策略使用各子数据集的采样概率长度必须等于数据集个数seedint / NoneNonenumpy RNG 种子用于复现采样序列global_rankint0当前进程的全局 rank用于 map 风格数据集的分片world_sizeint1总进程数用于 map 风格数据集的分片2.2 三种采样策略的实现原理在构造函数中ConcatDataset根据sampling_technique绑定对应的静态生成器见 dataset.py#L74-L871温度采样temperature_generator见 dataset.py#L166-L186这是默认策略。核心是先按各子数据集长度计算初始概率再施加温度缩放p np.array(lengths) / np.sum(lengths) # 按数据量占比 p np.power(p, 1 / temp) # 温度缩放 p p / np.sum(p) # 重新归一化其效果是sampling_temperature 1时等价于按数据量比例采样大库占绝对主导温度越大如默认值 5分布越趋向均匀小数据集获得更多被采到的机会温度越小趋向 0分布越尖锐几乎总是采样最大的数据集。2随机采样random_generator见 dataset.py#L196-L211直接使用用户显式传入的sampling_probabilities作为多项分布参数通过np_rng.choice(np.arange(num), pp)每次抽取一个数据集 ID。若未传sampling_probabilities则退化为各数据集等概率1 / len(datasets)。该策略不依赖数据集规模适合按业务需求硬编码采样比例的场景。3轮询采样round_robin_generator见 dataset.py#L188-L194最朴素的策略按0, 1, 2, ..., n-1, 0, 1, ...的顺序循环选择数据集保证每个数据集被公平轮询常用于调试或希望每个数据集被均匀覆盖的场景。2.3 底层迭代机制与 DataLoader 多进程交互ConcatDataset.__iter__见 dataset.py#L119-L161是理解其性能与正确性的关键探测 worker 环境通过torch.utils.data.get_worker_info()获取当前 DataLoader worker 的id与num_workers。若在主进程中worker_info is None则wid 0, wnum 1否则按len(range(wid, self.length, wnum))计算本 worker 应产出的样本数从而把总样本量均匀分摊到各 worker构造分片为每个子数据集基于global_rank/world_size切出本进程负责的索引区间并用torch.utils.data.Subset包装详见下文 2.4启动采样循环外层 while 循环从生成器中取数据集 IDind再从对应子数据集的迭代器取一个样本若某子数据集迭代器耗尽则调用get_iterable重新生成map 风格会重新洗牌索引并n - 1补偿该轮从而在单个 epoch 内也能无缝续接耗尽的子数据集长度约定__len__返回self.length——map 风格子数据集按len(dataset) // world_size累加可迭代子数据集直接累加len(dataset)最终乘以sampling_scale。由于__iter__内对每个子数据集先洗牌再消费且get_iterable见 dataset.py#L109-L117对 map 风格数据集使用np.random.shuffle(indices)生成打乱的索引迭代器因此ConcatDataset天然支持 shuffle 且可被重复迭代每次 epoch 都会重新洗牌。2.4 分布式训练下的数据分片对于 map 风格的子数据集ConcatDataset在迭代时按global_rank/world_size做进程级分片见 dataset.py#L130-L138start_idx (len(self.datasets[idx]) // self.world_size) * self.global_rank end_idx start_idx (len(self.datasets[idx]) // self.world_size) if self.global_rank self.world_size - 1: end_idx len(self.datasets[idx]) # 最后一个 rank 拿走余数部分 indices range(start_idx wid, end_idx, wnum) datasets.append(pt_data.Subset(self.datasets[idx], indices))也就是说每个 rank 只看到整个数据集的1 / world_size切片末尾 rank 吸收整除余数切片内再按 DataLoader 的 worker id 与 worker 数做二次切分。这保证了多卡训练时各进程数据不重叠、覆盖完整与单元测试 tests/collections/common/data/test_dataset_concat.py 中test_ranks_get_disjoint_shards的断言完全一致。三、ConcatMapDatasetMap 风格的多数据集采样ConcatMapDataset见 dataset.py#L214继承自torch.utils.data.Dataset语义上解决同样的问题但采用预构建索引表 __getitem__的实现方式与ConcatDataset形成互补。3.1 与 ConcatDataset 的差异维度ConcatDatasetConcatMapDataset基类IterableDataset惰性流式Dataset随机访问样本产出方式__iter__边采样边产出构造时一次性构建self.indices索引表__getitem__按索引返回采样序列确定性迭代时逐步决定构造完成后序列即固定受 seed 控制分片支持内置global_rank/world_size/worker 分片需由外部 DataLoader 或调用方配合分片ConcatMapDataset的构造参数更精简datasets、sampling_technique默认temperature、sampling_temperature默认 5、sampling_probabilities、seed没有shuffle、sampling_scale、global_rank、world_size等参数见 dataset.py#L229-L236。3.2 预构建索引与置换机制构造阶段见 dataset.py#L245-L298的核心逻辑为每个子数据集生成一份随机置换np_rng.permutation(len(x))作为该数据集的消费顺序若为round-robin总长度取max(self.lengths) * len(self.datasets)按np.arange(total_length) % len(datasets)循环分配数据集 ID保证每个数据集都被轮询到长数据集会被反复轮询若为random校验sampling_probabilities长度与数据集个数一致否则抛ValueError直接用归一化后的p做np_rng.choice抽取数据集 ID若为temperature同样按p lengths**(1/T)归一化后做加权抽取消费某数据集的置换序列时若耗尽则重新生成一份新置换并重置位置dataset_positions清零同时将该数据集标记为已耗尽循环终止条件是所有数据集都至少被耗尽过一次以此保证小数据集不会被饿死最终self.indices是一个(dataset_id, dataset_index)元组列表__getitem__(idx)见 dataset.py#L303-L305据此索引返回样本。因此ConcatMapDataset非常适合需要精确控制 epoch 长度、或希望在采样序列上与 shuffle 独立管理的场景——整个 epoch 的采样轨迹在构造时就被确定下来可复现性更强。四、在 ASR 数据管线中的实际接线ConcatDataset是 NeMo ASR 数据管线的底层基础设施被多个工厂函数直接调用。以 nemo/collections/asr/data/audio_to_text_dataset.py 为例1get_concat_char_dataset见 audio_to_text_dataset.py#L84-L130当配置中manifest_filepath是字符串列表时它会为每个 manifest 分别构造一个AudioToCharDataset再统一包进ConcatDataset。值得注意的细节get_concat_char_dataset支持传入形如[[dataset1, dataset2]]的额外嵌套层验证集场景下 ModelPT 会引入内部会先做一层扁平化处理。构造ConcatDataset时各采样参数均从 config 中读取并带有默认值dataset ConcatDataset( datasets, sampling_techniqueconfig.get(concat_sampling_technique, temperature), sampling_temperatureconfig.get(concat_sampling_temperature, 5), sampling_scaleconfig.get(concat_sampling_scale, 1), sampling_probabilitiesconfig.get(concat_sampling_probabilities, None), shuffleconfig.get(concat_shuffle, True), seedconfig.get(concat_sampling_seed, None), global_rankglobal_rank, world_sizeworld_size, )2get_concat_bpe_dataset见 audio_to_text_dataset.py#L167-L213与 char 版逻辑完全对称只是每个子数据集是AudioToBPEDataset采样参数同样通过concat_*前缀的配置项注入。3get_concat_tarred_dataset见 audio_to_text_dataset.py#L247-L291用于拼接多个tarred音频数据集每个数据集含tarred_audio_filepaths与对应 manifest这是大规模预训练中最常见的用法——多个 tar 分片集合作混合采样。由于 tarred 数据集是IterableDatasetConcatDataset会走kind iterable分支直接复用其迭代器而不做索引洗牌与 rank 分片。4语音标签任务nemo/collections/asr/data/audio_to_label_dataset.py#L159-L168 中的get_concat_tarred_speech_label_dataset以相同模式拼接多个 tarred 语音标签数据集如 VAD、说话人分类任务。5自监督预训练nemo/collections/asr/data/ssl_dataset.py 也导入了ConcatDataset用于自监督语音预训练的多语料混合。五、真实配置文件示例在 examples/asr/conf/ssl/nest/nest_fast-conformer.yaml 中可以看到concat_*配置项在自监督预训练SSL NEST FastConformer中的真实写法model: sample_rate: 16000 ... train_ds: manifest_filepath: ??? # 训练 manifest可为单个字符串或字符串列表 sample_rate: ${model.sample_rate} batch_size: 8 shuffle: true num_workers: 8 max_duration: 60.0 min_duration: 1.0 drop_last: true is_concat: false concat_sampling_technique: temperature # 混合采样策略temperature / random / round-robin concat_sampling_temperature: 1.0 # 温度1 时按数据量比例采样 is_tarred: false tarred_audio_filepaths: null shuffle_n: 2048 bucketing_strategy: synced_randomized ...这里concat_sampling_temperature: 1.0表示各 manifest 按原始数据量占比采样若希望小数据集获得更多权重可适当调大温度如 5 或更高。使用流程为将manifest_filepath配置为多个 manifest 的列表 → 保持is_concat相关语义数据工厂函数会依据 manifest 数量自动选择拼接路径→ 通过concat_*键控制采样行为。其余 ASR 训练/微调配置如 examples/asr/conf 下各模型 config 的train_ds段同样遵循这一约定。六、单元测试验证仓库在 tests/collections/common/data/test_dataset_concat.py 中为ConcatDataset提供了系统的单元测试是理解其行为契约的最佳范本test_iter_is_repeatable_across_epochs以round-robinshuffleFalseworld_size2构造两个数据集10 条与 7 条断言两个 epoch 的迭代结果完全一致且不修改原始数据集对象——验证了迭代的可重复性test_rank_shards_match_expected_slices构造global_rank0/1两个实例分别断言其产出的样本切片验证 rank 分片公式的正确性test_ranks_get_disjoint_shards断言两个 rank 的采样集合互不相交且并集等于全集——这是多卡训练无重叠、无遗漏的硬保证test_world_size_one_is_stable_and_complete单进程场景下验证轮询到长数据集被耗尽后短数据集会以重新洗牌的方式被再次轮询样本b0在尾部再次出现印证 2.3 节的耗尽续接机制test_temperature_weights_come_from_rank_shards验证温度采样的权重来自 rank 分片后的长度而非原始全长——即p shard_len^(1/T)归一化并断言实际采样频率与该权重吻合且比用未分片长度计算出的权重更贴近实测分布。七、选型与最佳实践综合以上源码分析给出如下选型建议默认首选ConcatDatasettemperature它同时处理了 map 风格数据集的 shuffle、rank/worker 分片与可迭代数据集的流式消费是 NeMo 数据管线默认路径且 ASR 工厂函数已内置concat_*配置映射接入成本最低显式控制各数据源比例用randomsampling_probabilities当你有明确的领域配比需求如 7:2:1时直接给出概率列表最直观注意概率列表长度必须与数据集个数一致追求采样序列可复现、可精确控制 epoch 长度用ConcatMapDataset其索引表在构造时一次性确定配合seed可完全复现适合评估/验证等对顺序敏感的流程调试与均衡覆盖用round-robin保证每个子数据集都被循环采到但会牺牲按规模加权的自然分布小数据集防饿死无论是调高sampling_temperature、使用round-robin还是利用ConcatMapDataset的全部耗尽才终止机制都能避免小数据集在长 epoch 中被边缘化分布式训练务必传入正确的global_rank/world_sizeConcatDataset依赖它们做进程级分片传入错误值会导致样本重叠或丢失可对照 test_dataset_concat.py 中的分片断言自查。总结ConcatDataset与ConcatMapDataset是 NeMo 多数据集混合训练的核心基础设施前者以可迭代方式提供了 temperature / random / round-robin 三种采样策略并内置了 DataLoader 多进程与多卡分片逻辑后者以预构建索引表的方式实现了可复现的 map 风格采样。二者在 ASR 字符/BPE/tarred 数据管线与自监督预训练配置中均有真实接线配合 test_dataset_concat.py 的测试契约开发者可以在自己的训练配置中放心使用concat_sampling_technique、concat_sampling_temperature、concat_sampling_probabilities等参数精确控制数据混合行为。【免费下载链接】SpeechA scalable generative AI framework built for researchers and developers working on Large Language Models, Multimodal, and Speech AI (Automatic Speech Recognition and Text-to-Speech)项目地址: https://gitcode.com/GitHub_Trending/nem/Speech创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考