新闻详情

新闻详情

首页 / 资讯中心 / 详情

Ray Data 加权数据集混合(Weighted Dataset Mixing)实战指南:从 per-block 到 random mixing 的训练数据配比控制

发布时间:2026/9/19 6:41:25来源:尧图网络
Ray Data 加权数据集混合(Weighted Dataset Mixing)实战指南:从 per-block 到 random mixing 的训练数据配比控制
Ray Data 加权数据集混合Weighted Dataset Mixing实战指南从 per-block 到 random mixing 的训练数据配比控制【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址: https://gitcode.com/gh_mirrors/ra/ray导读Ray Data 的Dataset.mix提供了一种流式streaming加权交错机制可以把多个数据集的 block 按目标行数比例交错成一个统一输出流用于解决类别/场景不平衡、多任务预训练配比、防灾难性遗忘等典型问题。本文以 Ray Data 官方文档 为骨架深入 Ray 仓库源码与测试用例完整讲解mix的 API 语义、per-block mixing 与 random mixing 两种策略、输入 block 尺寸标准化技巧、停止条件与当前已知限制帮助读者在 Ray Train 分布式训练中稳定控制每个全局 batch 的数据来源构成。一、为什么要做加权数据集混合在真实的机器学习训练流程中不同来源的数据往往有着截然不同的稀缺程度与训练价值直接将它们拼接union成一份数据集会导致训练分布失衡。mix面向的场景包括类别 / 场景平衡Class / scenario balancing对稀有场景或高难度任务做上采样让训练 batch 更频繁地看到它们。多任务预训练Multi-task pretraining以固定比例混合代码语料与网页文本语料。防灾难性遗忘Catastrophic forgetting prevention在训练新数据的同时保留一小部分旧数据集参与混合避免模型遗忘已学能力。与union只做简单行拼接不同见 dataset.py 中的 union 实现mix是一个流式交错算子它记录每个输入数据集已输出的累计行数并始终从落后于目标比例最多的那个数据集拉取下一个 block。从源码注释看mix_operator.py 的_select_most_behind_input每个输出 block 恰好来自某一个输入数据集但整体行数比例会逐渐收敛到weights指定的目标。Quickstart一个可直接运行的示例以下示例来自文档的 Quickstart展示了读取两个模拟数据集、按[0.75, 0.25]混合并交给 4 个 worker 的TorchTrainer消费的全过程import ray.data from ray.train.torch import TorchTrainer from ray.train import ScalingConfig # 分别读取并预处理各个数据源。 # 注意这里用 mock 数据做演示。 def preprocess(row): return row ds1 ray.data.from_items([{x: 1} for _ in range(750)]).map(preprocess) ds2 ray.data.from_items([{x: 2} for _ in range(250)]).map(preprocess) # 输出 batch 中约 75% 的行来自 ds125% 来自 ds2期望意义下。 mixed ds1.mix(ds2, weights[0.75, 0.25]) def train_fn_per_worker(config): shard ray.train.get_dataset_shard(train) for batch in shard.iter_torch_batches(batch_size128): print(batch) trainer TorchTrainer( train_loop_per_workertrain_fn_per_worker, scaling_configScalingConfig(num_workers4), datasets{train: mixed}, )要点拆解每个数据源独立完成读取与预处理如map(preprocess)再交给mix。weights列表的第一项对应self即调用mix的ds1后续项依次对应*other即ds2。混合后的mixed是一个新的Dataset可以直接作为 Ray Train 的datasets输入在train_fn_per_worker中通过ray.train.get_dataset_shard(train)分片读取。iter_torch_batches会把每个分片按batch_size128产出 PyTorch 兼容的 batch在数据并行训练中各 worker 的本地 batch 再聚合为全局 batch。二、mixing 策略总览文档把混合策略按混合粒度分为两类二者可以组合使用策略方式特点Per-block mixing仅调用mix每个输出 block 完整来自单一数据集block 间按目标比例交错Random mixingmix之后再跟一次 shuffle行被重新分布到 block 边界之外单个 batch 内直接混合多个数据源关键区别在于per-block mixing 的比例保证体现在 block 层面而 random mixing 通过洗牌把比例保证下沉到 batch 层面。下面分别展开。三、Per-block mixing按 block 粒度交错3.1 工作原理默认情况下每个输出 block 恰好来自一个输入数据集。mix内部维护每个数据源的累计行数每一步都从相对目标比例落后最多的那个数据集取下一个 block源码见 mix_operator.py 的_select_most_behind_inputgap self._weights[i] * total - self._rows_seen[i]gap越大表示该数据源越欠账。当 gap 相同时优先选择权重更高的数据源tie-break by weight再按输入序号决定。_try_output的逻辑还保证了确定性即使某个被选中的数据源的 block 尚未就绪也会等待而不是改从其他数据源取数mix_operator.py_try_output因此输出顺序不依赖 block 到达的时序。最终累计行数会收敛到目标比例。在文档的示例中ds1与ds2以[0.75, 0.25]混合且双方 block 等大输出流按 3:1 的模式交错随后被切分到 4 个训练 worker 上拼成全局 batch3.2 均匀 block 下的比例精确度当输入 block 大小均匀时比例在任意1 / min(weights)个 block 的窗口内都是精确的。例如weights[0.9, 0.1]可以保证每 10 个 block 的窗口内至少出现一次来自第二个数据集的 block。从测试用例test_mix_uneven_weights可以印证这一保证75/25 权重配合均匀 block每 4 个 block 内即形成正确比例test_mix.pytest_mix_three_datasets则验证了 50/30/20 权重在每 10 个 block 内形成正确比例test_mix.py。3.3 需要注意block 并不等于训练 batchRay Data 中 Blocksblock是数据传输单元它和训练 batch 并非一一对应。worker 构造一个 batch 时会从一个或多个 block 中拉取行。因此在 per-block mixing 下单个本地 batch 可能只包含一个数据源的行也可能包含多个数据源的行具体取决于 block 大小与 batch 大小的相对关系。这也引出了下一节介绍的 block 尺寸标准化方法。3.4 进阶标准化输入 block 尺寸如果输入数据集的 block 大小差异悬殊一个超大的 block 会暂时把该数据源推到目标比例之前虽然mix会在后续的拉取中自我纠正期望意义下比例依然正确但由少数几个 block 拼出的全局 batch 可能明显偏离目标比例。文档给出的解法是在mix之前用ds.repartition(target_num_rows_per_block)统一各数据源的 block 尺寸LOCAL_BATCH_SIZE 128 ds1 ray.data.from_items([{x: 1} for _ in range(750)]).map(preprocess) ds2 ray.data.from_items([{x: 2} for _ in range(250)]).map(preprocess) # 标准化 block 尺寸让比例在更紧的窗口内成立。 ds1 ds1.repartition(target_num_rows_per_blockLOCAL_BATCH_SIZE) ds2 ds2.repartition(target_num_rows_per_blockLOCAL_BATCH_SIZE) mixed ds1.mix(ds2, weights[0.75, 0.25])文档进一步建议如果行本身按字节计算很小可以考虑 repartition 到 batch size 的整数倍例如N * LOCAL_BATCH_SIZE避免把 block 切得过碎、徒增调度与传输开销。repartition的完整签名与参数说明可参考 ray.data.Dataset.repartition。四、Random mixing在 batch 内部混合多数据源4.1 为什么还需要 random mixingper-block mixing 的 batch 级比例质量取决于两个因素输入 block 的大小上一节已解决参与每个全局 batch 的 worker 数量。一个全局 batch 聚合了num_workers * grad_accum_steps个本地 batch而每个本地 batch 又来自某一个数据集因此每个全局 batch 中包含的本地 batch 数越多实际比例越贴近目标。极端情况下——单 worker、无梯度累积——每个全局 batch 就是本地 batch于是每个 batch 都只来自单一数据集。解决思路是在mix之后追加一次流式 shuffleshuffle 会把行重新分布到 block 边界之外使每个 batch 直接包含多个数据源的行并大致保持目标比例与训练 worker 数量无关。mix仍然负责比例控制shuffle 只是把比例摊薄到每个 batch 内部。4.2 Ray Data 中两种流式友好的 shuffle 方式文档推荐了两种不打断流式管线的 shuffleLocal buffer shuffle在DataIterator.iter_batches中通过local_shuffle_buffer_size参数启用map_batches shuffle自定义 shuffle 函数配合map_batches使用。文档给出的map_batches方案配合 NumPy 与 PyArrowimport numpy as np import pyarrow as pa LOCAL_BATCH_SIZE 128 ds1 ray.data.from_items([{x: 1} for _ in range(750)]).map(preprocess) ds2 ray.data.from_items([{x: 2} for _ in range(250)]).map(preprocess) ds1 ds1.repartition(target_num_rows_per_blockLOCAL_BATCH_SIZE) ds2 ds2.repartition(target_num_rows_per_blockLOCAL_BATCH_SIZE) mixed ds1.mix(ds2, weights[0.75, 0.25]) # 在 mix() 之后追加 shuffle实现 random mixing。 def random_shuffle(batch: pa.Table) - pa.Table: indices np.random.permutation(len(batch)) return batch.take(indices) # shuffle buffer 需要足够大才能保证跨数据集的混合质量。 SHUFFLE_BUFFER_SIZE 64 * LOCAL_BATCH_SIZE mixed mixed.map_batches(random_shuffle, batch_sizeSHUFFLE_BUFFER_SIZE, batch_formatpyarrow)两个实践要点batch_sizeSHUFFLE_BUFFER_SIZE文档建议把 shuffle buffer 设成64 * LOCAL_BATCH_SIZE量级buffer 越大跨数据集混合的随机性与比例质量越好。batch_formatpyarrow自定义 shuffle 函数接收pa.Table用np.random.permutation生成打乱索引后通过batch.take(indices)重排写法简洁且对 Arrow 格式完全原生。关于 shuffle 的更多细节local shuffle buffer 与 map_batches shuffle 的取舍可进一步参考 Ray Data 的 shuffling 解决方案。五、停止条件Stopping conditionsmix的第二个参数stopping_condition控制管线何时终止支持两种枚举值条件行为STOP_ON_LONGEST_DROP默认管线在最长数据集耗尽时结束。较短的数据集耗尽后自动掉队退出剩余 batch 只由仍在产出的数据集提供STOP_ON_SHORTEST管线在最短数据集耗尽时立即结束其他数据集被截断对应实现可以在 mix_operator.py 的_try_output中看到STOP_ON_SHORTEST下只要任一输入耗尽就整体停止STOP_ON_LONGEST_DROP下则跳过已耗尽的输入继续从剩余输入取数。此外n_ary_operator.py 的estimate_num_mix_outputs会按停止条件估算输出行数/block 数STOP_ON_SHORTEST下按weight * min(count_i / weight_i)截断计算STOP_ON_LONGEST_DROP下则取最长数据集的行数。注意STOP_ON_SHORTEST无法精确预估输出 block 数因为权重控制的是行比例而非 block 比例见 mix_operator.pynum_outputs_total。测试用例对两种停止条件都有覆盖test_mix_stop_on_shortest验证最短数据集耗尽即停test_mix.pytest_mix_stop_on_longest_drop验证较短数据集掉队后其余数据继续输出test_mix.py。更细粒度的估算行为由estimate_num_mix_outputs的单元测试覆盖test_mix.py例如STOP_ON_LONGEST_DROP对[100, 200]行、[0.5, 0.5]权重会输出 200 行而STOP_ON_SHORTEST只输出 100 行。API 签名见 dataset.py 的mixdef mix( self, *other: Dataset, weights: Optional[List[float]] None, stopping_condition: MixStoppingCondition MixStoppingCondition.STOP_ON_LONGEST_DROP, ) - Dataset值得注意的 API 细节weightsNone时默认等权[1.0] * len(datasets)。权重会被内部归一化不需要相加为 1。权重数量必须与数据集数量一致否则抛出ValueError。权重必须为正数n_ary_operator.py中校验any(weight 0)会报错。mix产出的数据集不可 lineage 序列化因此不能作为 Ray Tune 的可调超参数使用这一点与union相同源码中以.. caution::明确标注。六、当前限制与规避建议文档明确列出了三条限制使用时应格外注意避免在mix之后再做map/filter下游转换可能合并或拆分 block破坏mix提供的行比例保证。正确的做法是把每个数据集的转换全部放到mix之前执行。Schema 必须一致mix不会为你统一 schema。需要在混合前用map或select_columns让所有输入在结构上完全一致否则行为未定义。高度倾斜的权重当前限制目前所有输入数据集会并发执行且集群资源在它们之间均分。当权重严重倾斜例如[0.95, 0.05]时高权重数据集可能成为瓶颈低权重数据集却在空转。文档建议目前将权重控制在彼此 5 倍以内例如[0.4, 0.3, 0.2, 0.1]。从 mix_operator.py 的执行结构看MixOperator继承自NAryOperator通过_input_queues为每个输入维护独立缓冲、所有输入共享调度循环因此资源均分、并发执行是当前实现的固有特性重倾斜权重下的空闲问题属于实现层面的已知限制而非配置失误。七、使用建议小结结合文档与源码给出如下实践清单先各自预处理再混合read → map/filter/select_columns → mix所有 schema 统一与数据清洗都在mix之前完成。按需标准化 block需要更紧的 per-batch 比例窗口时在mix前用repartition(target_num_rows_per_blockN * LOCAL_BATCH_SIZE)统一输入 block 尺寸。追求 batch 内混合在mix后追加map_batchesshuffle或使用iter_batches(local_shuffle_buffer_size...)并让 shuffle buffer 足够大文档示例为64 * LOCAL_BATCH_SIZE。按数据长度选择停止条件想让长数据集完整参与训练选STOP_ON_LONGEST_DROP默认想让管线在最短数据集处整齐收尾选STOP_ON_SHORTEST。权重保持温和当前实现下尽量让权重比值不超过约 5 倍避免低权重数据源空转。不要混用 Tune 超参搜索mix结果不可序列化不能作为 Tune 的超参。延伸阅读使用 Ray Data Ray Train 进行分布式训练与数据注入data-ingest-torchRay Data 的 shuffling 解决方案ray.data.Dataset.repartition核心实现MixOperator、MixLogicalOperator测试用例python/ray/data/tests/test_mix.py【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址: https://gitcode.com/gh_mirrors/ra/ray创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

Unity半透明景深失效?双材质球方案彻底解决 2026/9/19 7:29:32

Unity半透明景深失效?双材质球方案彻底解决

1. 半透明景深失效的根源拆解半透明物体在Unity里做景深,十有八九会遇到一个很尴尬的现象:场景里那些玻璃、水面、粒子特效,在开了Post-Processing的Depth of Field之后,要么完全糊成一片,要么干脆像贴纸一样浮在画面上…

阅读更多 →
3D打印STL文件资源获取与优化指南 2026/9/19 7:29:32

3D打印STL文件资源获取与优化指南

1. STL文件资源获取全攻略作为一名玩了5年3D打印的老玩家,我深刻理解找到优质STL文件的重要性。刚开始接触3D打印时,我也经历过到处找模型却下载到各种问题文件的痛苦经历。今天就把我这些年积累的资源渠道和避坑经验全部分享给大家。STL文件是3D打印最常…

阅读更多 →
编译失败回退三次?Loop 工程让 Codex 走 TaoToken 查 2026/9/19 7:29:32

编译失败回退三次?Loop 工程让 Codex 走 TaoToken 查

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
FMEA打分指南:S/O/D评价准则与RPN优先级判定 2026/9/19 7:29:32

FMEA打分指南:S/O/D评价准则与RPN优先级判定

简介:《严重度发生度探测度评价准则》是一份面向质量工程师、产品设计及工艺开发人员的FMEA风险评估参考资料,适用于DFMEA和PFMEA过程中确定严重度、发生度、探测度三项评分,并据此计算RPN风险优先数。资源为单份PDF文档,约475KB&…

阅读更多 →
AIS数据采集与避碰系统实现:从解码、清洗到DCPA/TCPA预警 2026/9/19 7:29:32

AIS数据采集与避碰系统实现:从解码、清洗到DCPA/TCPA预警

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →
单目视觉定位实战:PNP算法与ArUco标记在机器人导航中的应用 2026/9/19 7:26:32

单目视觉定位实战:PNP算法与ArUco标记在机器人导航中的应用

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞