[Refactor] Add index-based data packing and worker-side materialization - #2055
[Refactor] Add index-based data packing and worker-side materialization#2055YanhuiDua wants to merge 4 commits into
Conversation
|
@claude review |
|
Claude: Summary本 PR 将 RL 训练的 pack 规划( ProduceBatchResult impact: not affected —— RoutedExperts impact: affected —— padding 侧 routed-experts 构造从 controller 移入 Ray concurrency impact: not affected —— Main Flowchart after this PRflowchart TD
A[BaseRLTrainer._prepare_train_data<br/>产出 WorkerInputItem<br/>advantages 改为 Tensor] --> B[TrainingController.fit]
B --> C[提取 data_lengths]
C --> D[RLDataPacker.pack<br/>仅索引规划 + padding 统计]
D --> E[ray.put 整个 data_batches<br/>嵌套 ObjectRef 下发]
E --> F[TrainingWorker.fit<br/>ray.get 全量 batch]
F --> G[_materialize_packs<br/>按索引选样本]
G --> H[_single_pack<br/>SequenceContext.cat + 拼接]
H --> I[_create_padding_item<br/>全 padding 分支 routed-experts 出错]
I --> J[按 packed_batch_num_per_step 执行 optimizer step]
B --> K[返回 TrainingLogInfo<br/>pack_time / train_time / padding_tokens]
style D fill:#cde4ff
style G fill:#cde4ff
style E fill:#ffe4b5
style I fill:#ffb3b3
核心原理实现与单测
抽象与信息隐藏评估
单测建议
其他 IssuesCritical
Warning
VerdictREQUEST_CHANGES |
| if pad_len == self.config.pack_max_length: | ||
| pad_rand_index = torch.randint(low=0, high=1, size=(1, 1, 1)) | ||
| else: | ||
| pad_rand_index = torch.randint(low=0, high=self._pack_n_routed_experts, size=(pad_len, 1, 1)) |
There was a problem hiding this comment.
Claude: [正确性] [复杂] 全 padding pack(pack_indices 为空,由 _align_pack_count 与 greedy 的 total_pack_indices.extend([[] ...]) 产生)仍沿用旧的 size=(1,1,1) dummy。
旧代码里这个 dummy 就是该 item 自身的 seq_ctx,_add_rollout_routed_experts 走 tensor 分支并扩展到 pack_max_length;现在 padding item 一定会经过 SequenceContext.cat(worker.py:640),而 cat 把 rollout_routed_experts 收集成 list(sequence_context.py:356-357),于是 fit 走 list 分支只生成 size=(1, num_hidden_layers, num_experts_per_tok),随后 worker.py:559-561 的 rollout_routed_experts.size(0) == input_ids.size(1) 断言必然失败,开启 routed-experts 的 MoE RL 训练会崩溃。
建议去掉该特判,统一按 pad_len 构造:
| if pad_len == self.config.pack_max_length: | |
| pad_rand_index = torch.randint(low=0, high=1, size=(1, 1, 1)) | |
| else: | |
| pad_rand_index = torch.randint(low=0, high=self._pack_n_routed_experts, size=(pad_len, 1, 1)) | |
| pad_rand_index = torch.randint(low=0, high=self._pack_n_routed_experts, size=(pad_len, 1, 1)) |
RoutedExperts impact: 全 padding pack 的 routed-experts 张量长度与 input_ids 不一致,训练直接中断。
| @@ -286,7 +285,7 @@ def test_rl_train_with_sft(self): | |||
| train_controller.onload(target="all") | |||
| log_infos = train_controller.fit(data_batches, pack_max_length=1024, rollout_idx=1) | |||
There was a problem hiding this comment.
Claude: [测试] 该测试的 train_worker_cfg 用的是 pack_max_length=2048(本文件 L135),但 L280/286/316 仍传 pack_max_length=1024,会被 controller 新增的一致性校验(controller.py:81-84)直接抛 ValueError,测试必然失败。
概述
本 PR 为 RL 训练引入了可配置的数据 Pack 能力,并将 Pack 规划与实际 Tensor 拼接进行了解耦。
RLDataPacker现在只生成如下结构的索引规划:[dp_rank][optimizer_step][pack][sample_index]。实际的数据获取、Tensor 拼接和 Padding 操作由各个 TrainingWorker 在本地完成。主要改动
数据流程
BaseRLTrainer
│
├─ 生成 WorkerInputItem 列表
│
▼
TrainingController
├─ 提取样本长度
├─ 生成仅包含索引的 Pack Plan
└─ 将原始 Batch 写入 Ray Object Store
│
▼
TrainingWorker
├─ 获取原始 Batch
├─ 根据分配到的索引选择样本
├─ 执行实际 Pack 和 Padding
└─ 按规划执行 Optimizer Step