From 60c40c9eec1bfa9811617deb8f17c9f7d61cc2b3 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Sun, 20 Sep 2026 11:56:16 +0000 Subject: [PATCH 1/4] [Refactor] Move train data preparation into AgentLoop and TrainingController - AgentLoop.canonicalize_train_fields (base + localhost/sandbox overrides) builds the unified full-sequence train fields (input_ids/labels/logprobs) at generation time; semantic holes are baked into labels by the loops. - RolloutState drops response_mask; labels become the only supervision carrier. agent_loop_type records the producing loop class name and AGENTIC_AGENT_LOOP_TYPES discriminates agentic full-sequence samples. - TrainingController.fit accepts list[list[RolloutState]] and absorbs validation, session-clustered advantages, shift/tensorization, seq_ctx, teacher fields and data_info; BaseRLTrainer._prepare_train_data is gone. - calculate_group_effective_response_masks now bakes token staleness into labels in place (monotone, convergent) and its agentic exclusion uses agent_loop_type instead of the input_ids/labels presence heuristic. --- docs/zh_cn/rl/advanced_tutorial/agent_loop.md | 42 +- docs/zh_cn/rl/advanced_tutorial/judger.md | 10 +- .../verl_agent/common/agent_loop_verl_tool.py | 13 +- .../rl/test_multi_task_agent_loop_manager.py | 45 +- tests/rl/test_prepare_train_data.py | 552 ++++++++++++++---- tests/rl/test_producer.py | 4 - tests/rl/test_replay_buffer.py | 3 - tests/rl/test_rl_colocate_trainer.py | 26 +- .../test_rl_colocate_trainer_integration.py | 56 +- tests/rl/test_rl_disaggregated_trainer.py | 10 +- tests/rl/test_staleness_policy.py | 98 +++- xtuner/v1/data_proto/rl_data.py | 136 ++++- xtuner/v1/rl/agent_loop/agent_loop.py | 70 ++- xtuner/v1/rl/agent_loop/gsm8k_with_tool.py | 13 +- .../agent_in_localhost_loop.py | 39 +- .../agent_in_sandbox_loop.py | 40 +- .../rl/agent_loop_manager/disagg_producer.py | 6 +- .../v1/rl/agent_loop_manager/produce_utils.py | 5 +- xtuner/v1/rl/agent_loop_manager/producer.py | 6 +- xtuner/v1/rl/replay_buffer.py | 2 +- xtuner/v1/rl/trainer/controller.py | 304 +++++++++- xtuner/v1/rl/trainer/worker.py | 21 +- xtuner/v1/train/rl_trainer.py | 419 +------------ 23 files changed, 1251 insertions(+), 669 deletions(-) diff --git a/docs/zh_cn/rl/advanced_tutorial/agent_loop.md b/docs/zh_cn/rl/advanced_tutorial/agent_loop.md index b72ddd4d39..e87980d240 100644 --- a/docs/zh_cn/rl/advanced_tutorial/agent_loop.md +++ b/docs/zh_cn/rl/advanced_tutorial/agent_loop.md @@ -84,16 +84,17 @@ AgentLoop 输入和输出都是 `RolloutState`。如果后续使用预置 `RLTra AgentLoop 返回的 `RolloutState` 如果要进入训练,至少需要满足: - `status == Status.COMPLETED`。预置 trainer 会跳过 `ABORTED`、`FILTERED`、`FAILED` 的样本组。 -- `response_ids` 非空。`_prepare_train_data()` 用它构造训练 token。 +- `response_ids` 非空。`canonicalize_train_fields()` 用它构造训练 token。 - `response` 非空。Judger 和轨迹保存依赖文本 response。 -- `reward["score"]` 存在。`_prepare_train_data()` 会直接读取它计算 advantage。 +- `reward["score"]` 存在。训练 controller 会直接读取它计算 advantage。 如果提供以下字段,还需要满足长度约定: - `logprobs`:长度必须等于 `len(response_ids)`。 -- `response_mask`:长度必须等于 `len(response_ids)`。mask 为 `0` 的 token 会被转成训练 label `-100`,对应 advantage 也会置为 `0.0`。 -这也是自定义 AgentLoop 最容易出错的地方:工具返回、环境反馈、系统插入内容等不是模型生成的 token,应该在 `response_mask` 中置为 `0`,并给对应 `logprobs` 填 `0.0`。 +labels 是唯一的监督载体。`AgentLoop.generate_group()` 会把产出 loop 的类名写入 `agent_loop_type`,并在末尾调用 `canonicalize_train_fields()`,为 prompt+response 型样本构造全序列 `input_ids`/`labels`/`logprobs`(三者等长、未 shift)。这也是自定义 AgentLoop 最容易出错的地方:工具返回、环境反馈、系统插入内容等不是模型生成的 token,不参与训练——对应 label 直接写 `-100`,`logprobs` 填 `0.0`。prompt+response 型样本不需要手动构造这些字段(基类会兜底);自行组装全序列的多轮 loop 则必须自己把语义洞烙进 labels。训练侧只做 shift 和 advantage 计算,不处理任何掩码。 + +agentic 全序列 loop(如 `AgentInLocalhostLoop`、`AgentInSandboxLoop`)的类名需要登记到 `xtuner.v1.data_proto.rl_data.AGENTIC_AGENT_LOOP_TYPES`,其样本才会被 token 级 staleness 排除;未登记的自定义 loop 默认按 prompt+response(reasoning)样本处理,`agent_loop_type` 为 `None`(未记录,如手工构造的样本)时同样按 prompt+response 处理。 ## SingleTurnAgentLoop @@ -147,7 +148,7 @@ agent_loop_config = SingleTurnAgentLoopConfig( 自定义 AgentLoop 通常需要做四件事: 1. 继承 `AgentLoop`,实现 `generate_sample()`。 -2. 在 `generate_sample()` 中维护 `tokens`、`sample_params`、`response_ids`、`response`、`logprobs`、`response_mask`、`status`,不要在这里调用外部 Judger。 +2. 在 `generate_sample()` 中维护 `tokens`、`sample_params`、`response_ids`、`response`、`logprobs`、`status`,不要在这里调用外部 Judger。prompt+response 型样本的训练字段(`input_ids`/`labels`)由基类 `canonicalize_train_fields()` 兜底构造;自行组装全序列的 loop 需自己写 `input_ids`/`labels`,并把语义洞直接烙进 labels。 3. 若覆盖 `generate_group()`,在其中显式编排 Teacher、Judger 和组级过滤;没有 validity check 时应让完成的样本尽早进入 Teacher,有 validity check 时应只把过滤通过的完整组发送给 Teacher。 4. 继承 `AgentLoopConfig`,实现 `build_local()`,这样才能接入 `TaskSpecConfig.agent_loop_config`,并复用 Ray actor 构建逻辑。 @@ -200,7 +201,7 @@ class ToolAgentLoop(AgentLoop): async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: final_response_ids: list[int] = [] final_logprobs: list[float] = [] - final_response_mask: list[int] = [] + final_supervised: list[int] = [] cur_tokens = list(rollout_state.tokens or rollout_state.prompt_ids or []) remaining_tokens = self.sample_params.max_tokens @@ -221,7 +222,7 @@ class ToolAgentLoop(AgentLoop): final_response_ids.extend(response_ids) final_logprobs.extend(logprobs) - final_response_mask.extend([1] * len(response_ids)) + final_supervised.extend([1] * len(response_ids)) cur_tokens.extend(response_ids) tool_tokens = self._run_tool_and_encode_result(rollout_state) @@ -230,7 +231,7 @@ class ToolAgentLoop(AgentLoop): final_response_ids.extend(tool_tokens) final_logprobs.extend([0.0] * len(tool_tokens)) - final_response_mask.extend([0] * len(tool_tokens)) + final_supervised.extend([0] * len(tool_tokens)) cur_tokens.extend(tool_tokens) remaining_tokens = self.sample_params.max_tokens - len(final_response_ids) @@ -239,11 +240,17 @@ class ToolAgentLoop(AgentLoop): rollout_state.response_ids = final_response_ids[: self.sample_params.max_tokens] rollout_state.logprobs = final_logprobs[: self.sample_params.max_tokens] - rollout_state.response_mask = final_response_mask[: self.sample_params.max_tokens] + supervised = final_supervised[: self.sample_params.max_tokens] rollout_state.response = self.tokenizer.decode(rollout_state.response_ids) - assert len(rollout_state.response_ids) == len(rollout_state.logprobs) - assert len(rollout_state.response_ids) == len(rollout_state.response_mask) + # 语义洞直接烙进 labels:工具/环境 token 的 label 为 -100,不参与训练。 + prompt_ids = list(rollout_state.prompt_ids or []) + rollout_state.input_ids = prompt_ids + list(rollout_state.response_ids) + rollout_state.labels = [-100] * len(prompt_ids) + [ + resp_id if flag else -100 for resp_id, flag in zip(rollout_state.response_ids, supervised) + ] + + assert len(rollout_state.input_ids) == len(rollout_state.labels) == len(rollout_state.logprobs) if rollout_state.status == Status.COMPLETED and self.judger is not None and not self.enable_batch_judge: rollout_state = await self.run_judger(rollout_state) @@ -252,8 +259,8 @@ class ToolAgentLoop(AgentLoop): 这个例子强调两个约定: -- 模型生成 token 的 `response_mask` 为 `1`。 -- 工具或环境插入 token 的 `response_mask` 为 `0`,`logprobs` 填 `0.0`。 +- 模型生成 token 的 label 保留 token id(参与训练)。 +- 工具或环境插入 token 的 label 写 `-100`,`logprobs` 填 `0.0`。 ### 覆盖 generate_group @@ -302,7 +309,7 @@ class CustomAgentLoop(AgentLoop): return rollout_state ``` -如果任务的多轮上下文、工具结果或环境状态不能用默认 handler 合并,需要自己定义续跑逻辑。核心原则是:续跑后的 `response_ids`、`response`、`logprobs`、`response_mask` 必须仍然表示完整 response,而不是只有本轮新增部分。 +如果任务的多轮上下文、工具结果或环境状态不能用默认 handler 合并,需要自己定义续跑逻辑。核心原则是:续跑后的 `response_ids`、`response`、`logprobs`(以及自组装 loop 的 `input_ids`/`labels`)必须仍然表示完整 response,而不是只有本轮新增部分。 ## 在训练配置中使用 @@ -341,8 +348,9 @@ agent_loop_manager_cfg = AgentLoopManagerConfig( - `generate_sample()` 是否只处理单条 `RolloutState`。 - 推理前是否设置了 `rollout_state.tokens`。 - 每次调用 `rollout_ctl.generate.remote()` 前是否设置了本轮 `sample_params`。 -- 返回训练前,`response_ids`、`response`、`logprobs`、`response_mask` 是否完整且长度一致。 -- 非模型生成 token 是否在 `response_mask` 中置为 `0`。 +- 返回训练前,`response_ids`、`response`、`logprobs` 是否完整且长度一致;自组装全序列的 loop 还需保证 `input_ids`/`labels` 齐全且等长。 +- 非模型生成 token 的 label 是否已写 `-100`(prompt+response 型样本由基类 canonicalize 兜底;自组装 loop 自行烙入)。 +- 若是 agentic 全序列 loop,类名是否已登记到 `AGENTIC_AGENT_LOOP_TYPES`。 - 需要 Judger 时,是否通过 `self.run_judger(...)` 调用打分,以复用 pause/cancel 处理。 -- 若使用 `_prepare_train_data()`,是否保证最终有 `reward["score"]`。 +- 是否保证最终有 `reward["score"]`。 - 若使用 async partial rollout,是否正确处理 `enable_partial_rollout` 和历史 response 合并。 diff --git a/docs/zh_cn/rl/advanced_tutorial/judger.md b/docs/zh_cn/rl/advanced_tutorial/judger.md index f4737a214b..0ec3cf0669 100644 --- a/docs/zh_cn/rl/advanced_tutorial/judger.md +++ b/docs/zh_cn/rl/advanced_tutorial/judger.md @@ -124,11 +124,11 @@ async def batch_judge(self, rollout_states: list[RolloutState]) -> list[RolloutS 在预置 `SingleTurnAgentLoop` 中,进入 Judger 前的 `RolloutState` 通常已经包含: -- `prompt_ids`:prompt token,后续 `_prepare_train_data()` 会用它拼接训练输入。 +- `prompt_ids`:prompt token,后续训练 controller 会用它拼接训练输入。 - `response_ids`:模型生成的 response token,后续训练数据直接依赖它。 - `response`:模型生成的文本,`NativeJudger` 会用它计算 reward。 - `logprobs`:response token 的 rollout logprob。如果存在,长度必须和 `response_ids` 一致。 -- `response_mask`:哪些 response token 参与训练。为空时,`_prepare_train_data()` 默认所有 response token 都参与训练。 +- `labels`:监督目标(labels 是唯一监督载体,语义上不参与训练的位置为 `-100`)。prompt+response 型样本由 loop 侧 `canonicalize_train_fields()` 兜底构造。 - `reward_model["ground_truth"]`:标准答案或标签,预置 rule-based Judger 通常依赖这个字段。 `NativeJudger` 传给 `reward_handler` 的字段更窄: @@ -154,13 +154,13 @@ rollout_state.reward = { } ``` -如果这批数据后续要进入 `_prepare_train_data()`,必须满足: +如果这批数据后续要进入训练 controller,必须满足: - `reward` 不能为 `None`。 -- `reward` 必须包含数值型 `score` 字段,因为 `_prepare_train_data()` 会直接读取 `data.reward["score"]` 来计算 advantage。 +- `reward` 必须包含数值型 `score` 字段,因为 controller 会直接读取 `data.reward["score"]` 来计算 advantage。 - `status` 不能是 `ABORTED`、`FILTERED` 或 `FAILED`。 - `response` 和 `response_ids` 都不能为空。 -- 如果提供 `response_mask`,长度必须和 `response_ids` 一致。mask 为 `0` 的 token 会在训练 label 中变成 `-100`,对应 advantage 也会置为 `0.0`。 +- labels 中 `-100` 的位置不参与训练,对应 advantage 也会置为 `0.0`。 - 如果提供 `logprobs`,长度必须和 `response_ids` 一致。 其他 reward 字段可以按任务需要扩展,例如 `acc`、`format`、`tool_ok`、`reason` 等。 diff --git a/recipe/verl_agent/common/agent_loop_verl_tool.py b/recipe/verl_agent/common/agent_loop_verl_tool.py index f78ef27264..759c5bce95 100644 --- a/recipe/verl_agent/common/agent_loop_verl_tool.py +++ b/recipe/verl_agent/common/agent_loop_verl_tool.py @@ -134,11 +134,18 @@ async def generate_sample(self, rollout_state: RolloutState) -> RolloutState: # TODO: handle samples with corrupted tool tokens ? # convert verl_tool_agent_loop output to rollout_state - rollout_state.prompt_ids = output.prompt_ids - rollout_state.response_ids = output.response_ids + # verl 的 response_mask(工具输出位为 0)是本 loop 的语义监督,直接烙进 labels(-100)。 + prompt_ids = list(output.prompt_ids) + response_ids = list(output.response_ids) + semantic_mask = output.response_mask if output.response_mask is not None else [1] * len(response_ids) + rollout_state.prompt_ids = prompt_ids + rollout_state.response_ids = response_ids rollout_state.logprobs = output.response_logprobs rollout_state.routed_experts = output.routed_experts - rollout_state.response_mask = output.response_mask + rollout_state.input_ids = prompt_ids + response_ids + rollout_state.labels = [-100] * len(prompt_ids) + [ + resp_id if mask_value else -100 for resp_id, mask_value in zip(response_ids, semantic_mask) + ] rollout_state.status = Status.COMPLETED rollout_state.extra_fields.update(output.extra_fields) # judger needs response in text format diff --git a/tests/rl/test_multi_task_agent_loop_manager.py b/tests/rl/test_multi_task_agent_loop_manager.py index 515e3cbf21..35cdf7f500 100644 --- a/tests/rl/test_multi_task_agent_loop_manager.py +++ b/tests/rl/test_multi_task_agent_loop_manager.py @@ -312,13 +312,15 @@ def test_trainer_rejects_conflicting_legacy_and_task_validity_checks(self): _migrate_legacy_is_valid_sample_fn(cfg, MagicMock()) async def test_take_train_batch_applies_token_staleness_mask(self): - # 启用 token staleness 时,公开 produce_batch 路径应返回最终 effective mask。 + # 启用 token staleness 时,公开 produce_batch 路径应把过期 token 的监督清出 labels(只清零不恢复)。 state = RolloutState( rollout_id=1, group_id=1, message=[{"role": "user", "content": "prompt"}], prompt_ids=[1, 2], response_ids=[3, 4], + input_ids=[1, 2, 3, 4], + labels=[-100, -100, 3, 4], response_model_steps=[0, 4], status=Status.COMPLETED, ) @@ -344,7 +346,7 @@ async def test_take_train_batch_applies_token_staleness_mask(self): result = await manager.produce_batch(batch_size=1, train_step=5, model_step=4) - self.assertEqual(result.rollout_states[0][0].response_mask, [0, 1]) + self.assertEqual(result.rollout_states[0][0].labels, [-100, -100, -100, 4]) self.assertEqual(replay_buffer.task_token_stale_threshold_calls, [{"task": 4}]) async def test_take_train_batch_skips_agentic_token_staleness_mask(self): @@ -354,7 +356,8 @@ async def test_take_train_batch_skips_agentic_token_staleness_mask(self): message=[{"role": "user", "content": "prompt"}], input_ids=[1, 2], labels=[-100, 2], - response_mask=[1], + logprobs=[0.0, -0.1], + agent_loop_type="AgentInLocalhostLoop", status=Status.COMPLETED, ) strategy = _FakeProduceStrategy(token_stale_threshold=4) @@ -378,7 +381,41 @@ async def test_take_train_batch_skips_agentic_token_staleness_mask(self): result = await manager.produce_batch(batch_size=1, train_step=5, model_step=4) - self.assertEqual(result.rollout_states[0][0].response_mask, [1]) + self.assertEqual(result.rollout_states[0][0].labels, [-100, 2]) + + async def test_take_train_batch_sync_path_leaves_labels_untouched(self): + # 同步路径(无 token staleness)零操作:labels 在生成期已是最终态(语义洞已烙入)。 + state = RolloutState( + rollout_id=1, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + prompt_ids=[1, 2], + response_ids=[3, 4, 5], + input_ids=[1, 2, 3, 4, 5], + labels=[-100, -100, 4, 5, -100], + status=Status.COMPLETED, + ) + manager = AgentLoopManager( + task_runners=[ + _TaskRunner( + task_name="task", + agent_loop=_fake_agent_loop(), + produce_strategy=_FakeProduceStrategy(), + sampler=_FakeSampler(), + weight=1.0, + order=0, + ) + ], + replay_buffer=_FakeReplayBuffer( + rollout_states_by_task={"task": [[state]]}, + leftover_counts={}, + ), + rollout_controller=_fake_rollout_controller(), + ) + + result = await manager.produce_batch(batch_size=1, train_step=5, model_step=4) + + self.assertEqual(result.rollout_states[0][0].labels, [-100, -100, 4, 5, -100]) async def test_produce_batch_allocates_by_weight_and_returns_task_sorted_results(self): # 共卡 produce_batch 按 task 权重分配 batch,并按 task 名稳定返回训练数据和 leftover 统计。 diff --git a/tests/rl/test_prepare_train_data.py b/tests/rl/test_prepare_train_data.py index 33282d741f..f9cf161a87 100644 --- a/tests/rl/test_prepare_train_data.py +++ b/tests/rl/test_prepare_train_data.py @@ -1,17 +1,28 @@ -"""RLTrainer._prepare_train_data 的 PR-fast contract 测试。 +"""RL 训练数据构造的 contract 测试(loop 侧 canonicalize + controller 侧转换)。 + +重构后训练数据构造分两段,本文件只测两段的公开纯逻辑,不启动 trainer、Ray worker、模型或 +rollout backend: + +- loop 侧 ``canonicalize_train_fields``:基类为 prompt+response 型样本构造统一全序列字段 + (``input_ids``/``labels``/``logprobs`` 等长、未 shift;``agent_loop_type`` 由 generate_group + 写入产出 loop 的类名);localhost/sandbox 各自覆写为 + trace 全序列字段的校验(失败置 FAILED,不上抛)。 +- controller 侧 ``TrainingController._convert_rollout_groups``:消费已 canonicalize 且 labels 已定稿 + (语义洞直接烙在 labels 中)的状态,完成 shift、张量化、组级 advantage、seq_ctx 构造与 + ``data_info`` 统计。 -本文件只测试训练数据构造的纯逻辑,不启动 trainer、Ray worker、模型或 rollout backend。 当前测试点: - 文本样本的 input_ids、shifted_labels、rollout_logprobs、advantage 布局。 - 同一个 prompt 下多个 response 各自使用对应 reward / advantage。 -- VLM 样本使用 train_prompt_ids,并保留 multimodal 训练字段。 +- VLM 样本 loop 侧使用 train_prompt_ids 构造;controller 侧保留 multimodal 字段与 3D 位置。 - VLM M-RoPE:get_train_seq_ctx 用 global amax 续写 response position(对齐 SFT get_rope_index_3)。 -- 无效 rollout group 会被跳过。 -- 缺失 reward、logprob/mask 长度不一致、pack_max_length 过小时 fail fast。 +- 无效 rollout group 会被跳过;缺 reward 在 task_adv_weight>0 时 fail fast,=0 时零 advantage 参训。 +- 监督语义:labels 是唯一监督载体(语义洞 = labels 中的 -100);controller 不做任何掩码处理。 - 多条样本 pack 成一条序列后,input_ids/shifted_labels/advantages/rollout_logprobs(及 teacher 字段) 与输入位置逐一对齐(regression: advantage 曾比 input_ids 长 1 导致 pack 后整体错位)。 """ +import sys import unittest from typing import cast from unittest.mock import MagicMock, patch @@ -19,12 +30,24 @@ import numpy as np import torch -from xtuner.v1.data_proto.rl_data import RolloutState, Status, TeacherTargets, reset_rollout_response +from xtuner.v1.data_proto.rl_data import ( + RolloutState, + Status, + TeacherTargets, +) from xtuner.v1.datasets.mllm_tokenize_fn.qwenvl_rope2d import get_rope_index_3 from xtuner.v1.rl.distillation import DistillationConfig, DistillationTrainerAdapter, RolloutTeacherConfig from xtuner.v1.rl.loss import DistillationLossConfig -from xtuner.v1.rl.trainer.controller import TrainingController -from xtuner.v1.train.rl_trainer import BaseRLTrainer, get_train_seq_ctx +from xtuner.v1.rl.trainer.controller import TrainingController, get_train_seq_ctx + +# localhost/sandbox loop 模块顶层导入 lagent;测试环境不依赖其真实实现。 +# rate_limiter 必须整路径 stub:父级 MagicMock 没有 __path__,子模块 from-import 无法透传。 +for _stub_name in ("lagent", "lagent.utils", "lagent.utils.rate_limiter"): + sys.modules.setdefault(_stub_name, MagicMock()) + +from xtuner.v1.rl.agent_loop.agent_loop import AgentLoop # noqa: E402 +from xtuner.v1.rl.agent_loop.localhost_agent_loop.agent_in_localhost_loop import AgentInLocalhostLoop # noqa: E402 +from xtuner.v1.rl.agent_loop.sandbox_agent_loop.agent_in_sandbox_loop import AgentInSandboxLoop # noqa: E402 class _FakeAdvantageEstimator: @@ -37,14 +60,278 @@ def compute(self, rewards_tensor, group): return torch.tensor(self.values[: len(group)], dtype=torch.float32) -class TestPrepareTrainData(unittest.TestCase): - def _build_trainer(self, advantages: list[float]): - trainer = BaseRLTrainer.__new__(BaseRLTrainer) - trainer._advantage_estimator = _FakeAdvantageEstimator(advantages) - trainer._distillation = DistillationTrainerAdapter(None) - trainer.tokenizer = MagicMock(return_value={"input_ids": torch.tensor([[999]])}) - trainer.logger = MagicMock() - return trainer +class _PromptResponseLoop(AgentLoop): + """最小具体子类:single-turn 等prompt+response loop 通过继承获得基类 canonicalize 实现。""" + + async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: + return rollout_state + + +class TestAgentLoopCanonicalizeTrainFields(unittest.TestCase): + """基类默认实现:为 prompt+response 型样本构造统一全序列字段。""" + + def _make_loop(self) -> AgentLoop: + loop = _PromptResponseLoop.__new__(_PromptResponseLoop) + loop.logger = MagicMock() + loop.tokenizer = MagicMock(return_value={"input_ids": torch.tensor([[999, 998]])}) + return loop + + def _state( + self, + *, + uid: int = 1, + prompt_ids: list[int] | None = None, + response_ids: list[int] | None = None, + response: str | None = "response", + logprobs: list[float] | None = None, + status: Status = Status.COMPLETED, + extra_fields: dict | None = None, + input_ids: list[int] | None = None, + labels: list[int] | None = None, + ) -> RolloutState: + return RolloutState( + rollout_id=uid, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + prompt_ids=prompt_ids if prompt_ids is not None else [10, 11, 12], + response=response, + response_ids=response_ids if response_ids is not None else [20, 21, 22], + logprobs=logprobs, + status=status, + finish_reason="stop" if status == Status.COMPLETED else "error", + extra_fields=extra_fields or {}, + input_ids=input_ids, + labels=labels, + ) + + def test_builds_full_sequence_fields(self): + # prompt+response 样本构造统一约定:三者等长,labels 全监督,logprobs 前补 0。 + loop = self._make_loop() + state = self._state(response_ids=[20, 21, 22], logprobs=[0.1, 0.2, 0.3]) + + returned = loop.canonicalize_train_fields([state]) + + self.assertIs(returned[0], state) + self.assertEqual(state.input_ids, [10, 11, 12, 20, 21, 22]) + self.assertEqual(state.labels, [-100, -100, -100, 20, 21, 22]) + self.assertEqual(state.logprobs, [0.0, 0.0, 0.0, 0.1, 0.2, 0.3]) + self.assertEqual(state.status, Status.COMPLETED) + + def test_tokenizer_fallback_when_response_ids_missing(self): + loop = self._make_loop() + state = self._state(response="ok") + state.response_ids = None + + loop.canonicalize_train_fields([state]) + + loop.tokenizer.assert_called_once_with("ok", return_tensors="pt") + self.assertEqual(state.response_ids, [999, 998]) + self.assertEqual(state.input_ids, [10, 11, 12, 999, 998]) + self.assertEqual(state.labels, [-100, -100, -100, 999, 998]) + + def test_flattens_tensor_response_ids(self): + # rollout backend 可能返回 Tensor 形态的 response_ids,构造前需 flatten。 + loop = self._make_loop() + state = self._state() + state.response_ids = torch.tensor([[20, 21, 22]]) + + loop.canonicalize_train_fields([state]) + + self.assertEqual(state.response_ids, [20, 21, 22]) + self.assertEqual(state.input_ids, [10, 11, 12, 20, 21, 22]) + + def test_vlm_uses_train_prompt_ids(self): + # VLM 分支用 extra_fields["train_prompt_ids"] 作为训练 prompt 构造全序列。 + loop = self._make_loop() + state = self._state(prompt_ids=[1], response_ids=[102, 103], extra_fields={"train_prompt_ids": [100, 101]}) + + loop.canonicalize_train_fields([state]) + + self.assertEqual(state.input_ids, [100, 101, 102, 103]) + self.assertEqual(state.labels, [-100, -100, 102, 103]) + + def test_logprobs_stay_none_when_missing(self): + loop = self._make_loop() + state = self._state(logprobs=None) + + loop.canonicalize_train_fields([state]) + + self.assertIsNone(state.logprobs) + + def test_agentic_state_with_input_ids_is_untouched(self): + # 已有 input_ids 的全序列样本(agentic loop 产出)满足约定,零改动。 + loop = self._make_loop() + state = self._state(input_ids=[30, 31, 40], labels=[-100, -100, 40], logprobs=[0.0, -0.1, -0.2]) + before = (list(state.input_ids), list(state.labels), list(state.logprobs)) + + loop.canonicalize_train_fields([state]) + + self.assertEqual((state.input_ids, state.labels, state.logprobs), before) + loop.tokenizer.assert_not_called() + + def test_non_completed_state_is_skipped(self): + loop = self._make_loop() + state = self._state(status=Status.FAILED) + + loop.canonicalize_train_fields([state]) + + self.assertIsNone(state.input_ids) + self.assertIsNone(state.labels) + self.assertEqual(state.status, Status.FAILED) + + def test_missing_prompt_ids_marks_sample_failed(self): + # 构造失败容错:单个样本置 FAILED,不上抛到 producer。 + loop = self._make_loop() + state = self._state() + state.prompt_ids = None + + loop.canonicalize_train_fields([state]) + + self.assertEqual(state.status, Status.FAILED) + self.assertIn("canonicalize_train_fields failed", state.error_msg) + loop.logger.error.assert_called_once() + + def test_failure_is_isolated_within_group(self): + loop = self._make_loop() + bad = self._state(uid=1) + bad.prompt_ids = None + good = self._state(uid=2) + + loop.canonicalize_train_fields([bad, good]) + + self.assertEqual(bad.status, Status.FAILED) + self.assertEqual(good.status, Status.COMPLETED) + self.assertEqual(good.input_ids, [10, 11, 12, 20, 21, 22]) + + +class TestAgenticLoopCanonicalizeTrainFields(unittest.TestCase): + """localhost/sandbox 各自的覆写:trace 全序列字段只做统一约定校验。""" + + AGENTIC_LOOP_CLASSES = (AgentInLocalhostLoop, AgentInSandboxLoop) + + def _make_loop(self, loop_cls): + loop = loop_cls.__new__(loop_cls) + loop.logger = MagicMock() + return loop + + def _agentic_state( + self, + *, + uid: int = 1, + status: Status = Status.COMPLETED, + input_ids: list[int] | None = None, + labels: list[int] | None = None, + logprobs: list[float] | None = None, + ) -> RolloutState: + return RolloutState( + rollout_id=uid, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + prompt_ids=[10, 11], + response="ok", + response_ids=[20, 21], + status=status, + finish_reason="stop", + input_ids=input_ids, + labels=labels, + logprobs=logprobs, + ) + + def test_valid_full_sequence_state_is_unchanged(self): + for loop_cls in self.AGENTIC_LOOP_CLASSES: + with self.subTest(loop_cls=loop_cls.__name__): + loop = self._make_loop(loop_cls) + state = self._agentic_state( + input_ids=[30, 31, 40, 41, 42], + labels=[-100, -100, 40, 41, 42], + logprobs=[0.0, -0.1, -0.2, -0.3, -0.4], + ) + + returned = loop.canonicalize_train_fields([state]) + + self.assertIs(returned[0], state) + self.assertEqual(state.status, Status.COMPLETED) + self.assertIsNone(state.error_msg) + loop.logger.error.assert_not_called() + + def test_labels_length_mismatch_marks_sample_failed(self): + for loop_cls in self.AGENTIC_LOOP_CLASSES: + with self.subTest(loop_cls=loop_cls.__name__): + loop = self._make_loop(loop_cls) + state = self._agentic_state(input_ids=[30, 31, 40, 41, 42], labels=[-100, -100, 40, 41]) + + loop.canonicalize_train_fields([state]) + + self.assertEqual(state.status, Status.FAILED) + self.assertIn("canonicalize_train_fields failed", state.error_msg) + loop.logger.error.assert_called_once() + + def test_missing_labels_marks_sample_failed(self): + for loop_cls in self.AGENTIC_LOOP_CLASSES: + with self.subTest(loop_cls=loop_cls.__name__): + loop = self._make_loop(loop_cls) + state = self._agentic_state(input_ids=[30, 31, 40, 41, 42], labels=None) + + loop.canonicalize_train_fields([state]) + + self.assertEqual(state.status, Status.FAILED) + self.assertIn("labels length mismatch", state.error_msg) + + def test_logprobs_length_mismatch_marks_sample_failed(self): + for loop_cls in self.AGENTIC_LOOP_CLASSES: + with self.subTest(loop_cls=loop_cls.__name__): + loop = self._make_loop(loop_cls) + state = self._agentic_state( + input_ids=[30, 31, 40, 41, 42], + labels=[-100, -100, 40, 41, 42], + logprobs=[0.0, -0.1, -0.2], + ) + + loop.canonicalize_train_fields([state]) + + self.assertEqual(state.status, Status.FAILED) + self.assertIn("logprobs length mismatch", state.error_msg) + + def test_state_without_input_ids_is_skipped(self): + # eval 态样本显式置空训练字段且 COMPLETED,不能被构造或校验。 + for loop_cls in self.AGENTIC_LOOP_CLASSES: + with self.subTest(loop_cls=loop_cls.__name__): + loop = self._make_loop(loop_cls) + state = self._agentic_state() + + loop.canonicalize_train_fields([state]) + + self.assertIsNone(state.input_ids) + self.assertEqual(state.status, Status.COMPLETED) + loop.logger.error.assert_not_called() + + def test_non_completed_state_is_skipped(self): + for loop_cls in self.AGENTIC_LOOP_CLASSES: + with self.subTest(loop_cls=loop_cls.__name__): + loop = self._make_loop(loop_cls) + state = self._agentic_state(status=Status.FAILED) + + loop.canonicalize_train_fields([state]) + + self.assertIsNone(state.input_ids) + self.assertEqual(state.status, Status.FAILED) + loop.logger.error.assert_not_called() + + +class TestConvertRolloutGroups(unittest.TestCase): + """TrainingController._convert_rollout_groups 合同:shift、advantage、张量与统计。""" + + def _build_controller(self, advantages: list[float], task_adv_weight: float = 1.0) -> TrainingController: + controller = TrainingController.__new__(TrainingController) + controller.advantage_estimator = _FakeAdvantageEstimator(advantages) + controller.task_adv_weight = task_adv_weight + controller.distillation = DistillationTrainerAdapter(None) + controller.logger = MagicMock() + return controller + + def _convert(self, controller, data_groups, pack_max_length=128): + with patch("xtuner.v1.rl.trainer.controller.XTUNER_DETERMINISTIC", True): + return controller._convert_rollout_groups(data_groups, pack_max_length) def _state( self, @@ -54,7 +341,7 @@ def _state( prompt_ids: list[int] | None = None, response_ids: list[int] | torch.Tensor | None = None, logprobs: list[float] | None = None, - response_mask: list[int] | None = None, + supervised_mask: list[int] | None = None, reward: dict | None = None, status: Status = Status.COMPLETED, response: str = "response", @@ -66,16 +353,19 @@ def _state( labels: list[int] | None = None, teacher_tokens: list[int] | list[list[int]] | None = None, teacher_logprobs: list[float] | list[list[float]] | None = None, + agent_loop_type: str | None = None, ) -> RolloutState: - return RolloutState( + resolved_prompt_ids = prompt_ids if prompt_ids is not None else [10, 11, 12] + resolved_response_ids = response_ids if response_ids is not None else [20, 21, 22] + state = RolloutState( rollout_id=uid, group_id=group_id, message=[{"role": "user", "content": f"prompt {group_id}"}], - prompt_ids=prompt_ids if prompt_ids is not None else [10, 11, 12], + prompt_ids=resolved_prompt_ids, response=response, - response_ids=response_ids if response_ids is not None else [20, 21, 22], + response_ids=resolved_response_ids, logprobs=logprobs, - response_mask=response_mask, + agent_loop_type=agent_loop_type, reward=reward if reward is not None else {"score": 1.0}, status=status, finish_reason="stop" if status == Status.COMPLETED else "error", @@ -95,14 +385,25 @@ def _state( else None ), ) - - def _prepare(self, trainer, data_groups, pack_max_length=128): - with patch("xtuner.v1.train.rl_trainer.XTUNER_DETERMINISTIC", True): - return trainer._prepare_train_data(data_groups, pack_max_length=pack_max_length) + if input_ids is not None: + # agentic 全序列样本由 trace store 直接提供字段,按原样返回。 + return state + # prompt+response 样本:模拟 loop 侧 canonicalize 后的最终形态(语义洞直接烙在 labels)。 + resp_ids = list(resolved_response_ids) + state.input_ids = list(resolved_prompt_ids) + resp_ids + if supervised_mask is None: + state.labels = [-100] * len(resolved_prompt_ids) + resp_ids + else: + state.labels = [-100] * len(resolved_prompt_ids) + [ + resp_id if flag else -100 for resp_id, flag in zip(resp_ids, supervised_mask) + ] + if state.logprobs is not None: + state.logprobs = [0.0] * len(resolved_prompt_ids) + list(state.logprobs) + return state @staticmethod - def _enable_rollout_distillation(trainer, loss_config: DistillationLossConfig) -> None: - trainer._distillation = DistillationTrainerAdapter( + def _enable_rollout_distillation(controller, loss_config: DistillationLossConfig) -> None: + controller.distillation = DistillationTrainerAdapter( DistillationConfig( loss_config=loss_config, teachers=[RolloutTeacherConfig(name="teacher", endpoints=["http://teacher"])], @@ -112,18 +413,18 @@ def _enable_rollout_distillation(trainer, loss_config: DistillationLossConfig) - def test_text_path_builds_shifted_training_tensors(self): # 文本主路径固定 token 布局:input_ids 去掉 response 最后一个 token,label/logprob 对齐预测位置。 - trainer = self._build_trainer([1.5]) + controller = self._build_controller([1.5]) routed_experts = np.array([[1, 2], [3, 4]]) state = self._state( prompt_ids=[10, 11, 12], response_ids=[20, 21, 22], logprobs=[0.1, 0.2, 0.3], - response_mask=[1, 0, 1], + supervised_mask=[1, 0, 1], reward={"score": 1.0}, routed_experts=routed_experts, ) - data_batches, info = self._prepare(trainer, [[state]]) + data_batches, info = self._convert(controller, [[state]]) self.assertEqual(len(data_batches), 1) batch = data_batches[0] @@ -143,29 +444,23 @@ def test_text_path_builds_shifted_training_tensors(self): self.assertEqual(info["response_len/mean"], 3.0) self.assertEqual(info["prompt_len/mean"], 3.0) - def test_rerolled_state_without_semantic_mask_uses_all_response_tokens(self): - trainer = self._build_trainer([1.0]) - state = reset_rollout_response(self._state(response_mask=[0, 1, 0])) - state.response = "rerolled response" - state.response_ids = [30, 31] - state.logprobs = [0.1, 0.2] - state.reward = {"score": 1.0} - state.status = Status.COMPLETED - state.finish_reason = "stop" + def test_controller_consumes_labels_as_final_supervision(self): + # controller 不处理任何掩码语义:labels 是什么就消费什么(语义洞已在生成期烙进 labels)。 + controller = self._build_controller([1.0]) + state = self._state(response_ids=[20, 21, 22], logprobs=[0.1, 0.2, 0.3], supervised_mask=[1, 0, 1]) - data_batches, _ = self._prepare(trainer, [[state]]) + data_batches, _ = self._convert(controller, [[state]]) - self.assertIsNone(state.response_mask) - self.assertEqual(data_batches[0]["shifted_labels"].tolist(), [[-100, -100, 30, 31]]) - self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 1.0, 1.0]) + self.assertEqual(data_batches[0]["shifted_labels"].tolist(), [[-100, -100, 20, -100, 22]]) + self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 1.0, 0.0, 1.0]) def test_multi_sample_group_uses_each_sample_reward_and_advantage(self): # 同一个 prompt 下的多个 response 要分别使用自己的 reward 和 advantage。 - trainer = self._build_trainer([1.5, -2.0]) + controller = self._build_controller([1.5, -2.0]) first = self._state(uid=1, response_ids=[20, 21], reward={"score": 3.0}) second = self._state(uid=2, response_ids=[30, 31], reward={"score": -1.0}) - data_batches, info = self._prepare(trainer, [[first, second]]) + data_batches, info = self._convert(controller, [[first, second]]) self.assertEqual(len(data_batches), 2) self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 1.5, 1.5]) @@ -176,17 +471,17 @@ def test_multi_sample_group_uses_each_sample_reward_and_advantage(self): self.assertEqual(info["rewards/mean"], 1.0) self.assertEqual(info["advantages/min"], -2.0) self.assertEqual(info["advantages/max"], 1.5) - self.assertEqual(trainer._advantage_estimator.calls[0][0].tolist(), [3.0, -1.0]) + self.assertEqual(controller.advantage_estimator.calls[0][0].tolist(), [3.0, -1.0]) def test_advantage_stats_count_only_loss_active_tokens(self): - # advantages/mean|min|max 只统计 loss-active token: prompt 占位与 mask=0 的 token 不参与。 - trainer = self._build_trainer([2.0, -1.0]) + # advantages/mean|min|max 只统计 loss-active token: prompt 占位与 labels=-100 的 token 不参与。 + controller = self._build_controller([2.0, -1.0]) plain = self._state( uid=1, prompt_ids=[10, 11, 12], response_ids=[20, 21, 22, 23], logprobs=[0.1, 0.2, 0.3, 0.4], - response_mask=[1, 0, 1, 0], + supervised_mask=[1, 0, 1, 0], reward={"score": 2.0}, ) agentic = self._state( @@ -195,31 +490,30 @@ def test_advantage_stats_count_only_loss_active_tokens(self): labels=[-100, -100, 40, -100, 42], logprobs=[0.0, -0.1, -0.2, -0.3, -0.4], reward={"score": -1.0}, + agent_loop_type="AgentInLocalhostLoop", ) - _, info = self._prepare(trainer, [[plain, agentic]]) + _, info = self._convert(controller, [[plain, agentic]]) - # plain: 2 个 mask!=0 token 计入 2.0; agentic: 2 个 label!=-100 token 计入 -1.0。 + # plain: 2 个 label!=-100 的 response token 计入 2.0; agentic: 同理计入 -1.0。 self.assertEqual(info["advantages/mean"], 0.5) self.assertEqual(info["advantages/min"], -1.0) self.assertEqual(info["advantages/max"], 2.0) - def test_vlm_path_uses_train_prompt_ids_and_preserves_multimodal_fields(self): - # VLM 分支使用 extra_fields["train_prompt_ids"] 作为训练 prompt,并把图像字段带进 SequenceContext。 - trainer = self._build_trainer([0.25]) + def test_vlm_fields_pass_through_with_3d_position_extension(self): + # VLM 样本保留 multimodal 字段;3D position_ids 按 len_response_ids 补 response 段。 + controller = self._build_controller([0.25]) pixel_values = np.ones((1, 2, 3), dtype=np.float32) image_grid_thw = np.array([[1, 2, 3]], dtype=np.int32) position_ids = np.array([[[0, 1]], [[0, 1]], [[0, 1]]], dtype=np.int64) state = self._state( - prompt_ids=[1], + prompt_ids=[100, 101], response_ids=[102, 103], - response_mask=[1, 1], position_ids=position_ids, mm_info={"pixel_values": pixel_values, "image_grid_thw": image_grid_thw}, - extra_fields={"train_prompt_ids": [100, 101]}, ) - data_batches, _ = self._prepare(trainer, [[state]]) + data_batches, _ = self._convert(controller, [[state]]) seq_ctx = data_batches[0]["seq_ctx"] self.assertEqual(seq_ctx.input_ids.tolist(), [[100, 101, 102]]) @@ -264,7 +558,7 @@ def test_get_train_seq_ctx_mrope_continues_from_global_amax(self): torch.testing.assert_close(rl_position_ids, sft_position_ids) def test_mixed_agentic_and_vlm_reasoning_use_3d_position_ids(self): - trainer = self._build_trainer([0.5]) + controller = self._build_controller([0.5]) reasoning_state = self._state( uid=1, group_id=1, @@ -278,9 +572,10 @@ def test_mixed_agentic_and_vlm_reasoning_use_3d_position_ids(self): input_ids=[30, 31, 40, 41, 42], labels=[-100, -100, 40, 41, 42], logprobs=[0.0, -0.1, -0.2, -0.3, -0.4], + agent_loop_type="AgentInLocalhostLoop", ) - data_batches, _ = self._prepare(trainer, [[reasoning_state], [agentic_state]]) + data_batches, _ = self._convert(controller, [[reasoning_state], [agentic_state]]) self.assertEqual(tuple(data_batches[0]["seq_ctx"].position_ids.shape), (3, 1, 5)) agentic_position_ids = data_batches[1]["seq_ctx"].position_ids @@ -301,8 +596,8 @@ def test_agentic_topk_targets_include_token_ids_and_logprobs(self): use_policy_gradient=False, top_k=2, ) - trainer = self._build_trainer([0.0]) - self._enable_rollout_distillation(trainer, loss_config) + controller = self._build_controller([0.0]) + self._enable_rollout_distillation(controller, loss_config) state = self._state( input_ids=[10, 11, 20, 21, 22], labels=[-100, -100, 20, -100, 22], @@ -310,9 +605,10 @@ def test_agentic_topk_targets_include_token_ids_and_logprobs(self): teacher_tokens=[[100, 101], [102, 103], [104, 105]], teacher_logprobs=[[-0.5, -0.6], [-0.7, -0.8], [-0.9, -1.0]], extra_fields={"origin_data_source": "agent_math"}, + agent_loop_type="AgentInLocalhostLoop", ) - data_batches, _ = self._prepare(trainer, [[state]]) + data_batches, _ = self._convert(controller, [[state]]) self.assertEqual(len(data_batches), 1) batch = data_batches[0] @@ -355,18 +651,18 @@ def test_plain_topk_targets_include_masked_response_rows(self): use_policy_gradient=False, top_k=2, ) - trainer = self._build_trainer([0.0]) - self._enable_rollout_distillation(trainer, loss_config) + controller = self._build_controller([0.0]) + self._enable_rollout_distillation(controller, loss_config) state = self._state( prompt_ids=[10, 11, 12], response_ids=[20, 21, 22], - response_mask=[0, 1, 1], + supervised_mask=[0, 1, 1], teacher_tokens=[[102, 103], [104, 105]], teacher_logprobs=[[-0.7, -0.8], [-0.9, -1.0]], extra_fields={"origin_data_source": "agent_math"}, ) - data_batches, _ = self._prepare(trainer, [[state]]) + data_batches, _ = self._convert(controller, [[state]]) batch = data_batches[0] self.assertEqual(batch["shifted_labels"].tolist(), [[-100, -100, -100, 21, 22]]) @@ -388,14 +684,14 @@ def test_sampled_token_targets_align_for_plain_and_agentic_rollouts(self): loss_mode="k1", use_policy_gradient=True, ) - trainer = self._build_trainer([0.0, 0.0]) - self._enable_rollout_distillation(trainer, loss_config) + controller = self._build_controller([0.0, 0.0]) + self._enable_rollout_distillation(controller, loss_config) plain_state = self._state( uid=1, group_id=1, prompt_ids=[10, 11, 12], response_ids=[20, 21, 22], - response_mask=[0, 1, 1], + supervised_mask=[0, 1, 1], teacher_tokens=[21, 22], teacher_logprobs=[-0.7, -0.9], extra_fields={"origin_data_source": "agent_math"}, @@ -409,9 +705,10 @@ def test_sampled_token_targets_align_for_plain_and_agentic_rollouts(self): teacher_tokens=[40, 41, 42], teacher_logprobs=[-1.1, -1.2, -1.3], extra_fields={"origin_data_source": "agent_math"}, + agent_loop_type="AgentInLocalhostLoop", ) - data_batches, _ = self._prepare(trainer, [[plain_state], [agentic_state]]) + data_batches, _ = self._convert(controller, [[plain_state], [agentic_state]]) self.assertEqual(len(data_batches), 2) torch.testing.assert_close( @@ -427,61 +724,74 @@ def test_sampled_token_targets_align_for_plain_and_agentic_rollouts(self): def test_invalid_group_is_skipped(self): # FAILED/FILTERED/ABORTED group 不能进入训练 batch,也不能贡献训练样本数。 - trainer = self._build_trainer([1.0]) + controller = self._build_controller([1.0]) valid = self._state(uid=1, response_ids=[20, 21], reward={"score": 2.0}) failed = self._state(uid=2, status=Status.FAILED, response_ids=[30, 31], reward={"score": 4.0}) - data_batches, info = self._prepare(trainer, [[valid], [failed]]) + data_batches, info = self._convert(controller, [[valid], [failed]]) self.assertEqual(len(data_batches), 1) self.assertEqual(info["training_samples"], 1) - trainer.logger.error.assert_called_once() + controller.logger.error.assert_called_once() def test_missing_reward_score_fails_fast(self): # reward 必须包含 score,否则 advantage 计算前后语义都不明确。 - trainer = self._build_trainer([1.0]) + controller = self._build_controller([1.0]) state = self._state(reward={"other": 1.0}) with self.assertRaisesRegex(ValueError, "missing.*score"): - self._prepare(trainer, [[state]]) + self._convert(controller, [[state]]) - def test_logprobs_must_match_response_ids_length(self): - # rollout logprobs 和 response_ids 必须逐 token 对齐。 - trainer = self._build_trainer([1.0]) - state = self._state(response_ids=[20, 21, 22], logprobs=[0.1, 0.2]) + def test_missing_reward_with_zero_task_adv_weight_trains_with_zero_advantage(self): + # task_adv_weight=0 时(纯 OPD)缺 reward 不崩溃:样本以零 advantage 参训,不计入 rewards 统计。 + controller = self._build_controller([1.0], task_adv_weight=0.0) + state = self._state() + state.reward = None - with self.assertRaises(AssertionError): - self._prepare(trainer, [[state]]) + data_batches, info = self._convert(controller, [[state]]) - def test_response_mask_must_match_response_ids_length(self): - # response_mask 参与 label 和 advantage mask,长度不一致时必须直接失败。 - trainer = self._build_trainer([1.0]) - state = self._state(response_ids=[20, 21, 22], response_mask=[1, 0]) + self.assertEqual(len(data_batches), 1) + self.assertEqual(data_batches[0]["advantage"], [0.0] * 5) + self.assertEqual(info["training_samples"], 1) + self.assertEqual(info["rewards/mean"], 0.0) + self.assertEqual(controller.advantage_estimator.calls, []) - with self.assertRaises(AssertionError): - self._prepare(trainer, [[state]]) + def test_group_with_mismatched_full_sequence_logprobs_is_skipped(self): + # 全序列 logprobs 与 input_ids 不对齐的样本由组校验拦截:整组跳过,不进训练。 + controller = self._build_controller([1.0]) + state = self._state( + input_ids=[30, 31, 40, 41, 42], + labels=[-100, -100, 40, 41, 42], + logprobs=[0.0, -0.1, -0.2, -0.3], + agent_loop_type="AgentInLocalhostLoop", + ) + + data_batches, info = self._convert(controller, [[state]]) + + self.assertEqual(data_batches, []) + self.assertEqual(info["training_samples"], 0) def test_input_ids_must_not_exceed_pack_max_length(self): # pack_max_length 过小时要在进入 packing 前失败,避免后续训练侧报错难定位。 - trainer = self._build_trainer([1.0]) + controller = self._build_controller([1.0]) state = self._state(prompt_ids=[10, 11, 12], response_ids=[20, 21, 22]) with self.assertRaises(AssertionError): - self._prepare(trainer, [[state]], pack_max_length=4) + self._convert(controller, [[state]], pack_max_length=4) -class TestPrepareTrainDataPackAlignment(unittest.TestCase): +class TestConvertRolloutGroupsPackAlignment(unittest.TestCase): """多条样本 pack 成一条序列后, pg_loss 消费前的张量逐位置对齐。 regression: advantage 曾按 `[adv]*len(prompt_ids) + [mask处理]` 构造, 比 input_ids 长 1, pack 拼接后 advantages 与 shifted_labels/rollout_logprobs 整体错位或形状不匹配。 """ - def _make_trainer(self, with_teacher_fields: bool = False) -> BaseRLTrainer: - trainer = BaseRLTrainer.__new__(BaseRLTrainer) - trainer.logger = MagicMock() + def _make_controller(self, with_teacher_fields: bool = False) -> TrainingController: + controller = TrainingController.__new__(TrainingController) + controller.logger = MagicMock() + controller.task_adv_weight = 1.0 distillation = MagicMock() - distillation.task_adv_weight = 1.0 if with_teacher_fields: # 与真实批次一致的逐位置 teacher 字段形状 (1, seq_len, k) / (1, seq_len)。 distillation.rollout_teacher_targets.side_effect = lambda state, shifted_labels: { @@ -492,32 +802,38 @@ def _make_trainer(self, with_teacher_fields: bool = False) -> BaseRLTrainer: else: distillation.rollout_teacher_targets.return_value = {} distillation.reward_scalars.return_value = {} - trainer._distillation = distillation - trainer._advantage_estimator = MagicMock() + controller.distillation = distillation + controller.advantage_estimator = MagicMock() # 确定性 advantage: advantage = reward - 0.5, 每个样本值唯一且 float32 精确。 - trainer._advantage_estimator.compute.side_effect = lambda rewards, representatives: rewards - 0.5 - return trainer + controller.advantage_estimator.compute.side_effect = lambda rewards, representatives: rewards - 0.5 + return controller def _make_sample( self, rollout_id: int, prompt_ids: list[int], response_ids: list[int], - response_mask: list[int], + supervised_mask: list[int], logprobs: list[float], reward: float, ) -> RolloutState: - return RolloutState( + state = RolloutState( rollout_id=rollout_id, message=[], prompt_ids=prompt_ids, response_ids=response_ids, - response_mask=response_mask, logprobs=logprobs, response="ok", reward={"score": reward}, status=Status.COMPLETED, ) + # 模拟 loop 侧 canonicalize 后的最终形态(语义洞直接烙在 labels)。 + state.input_ids = list(prompt_ids) + list(response_ids) + state.labels = [-100] * len(prompt_ids) + [ + resp_id if flag else -100 for resp_id, flag in zip(response_ids, supervised_mask) + ] + state.logprobs = [0.0] * len(prompt_ids) + list(logprobs) + return state @staticmethod def _expected_segment( @@ -526,15 +842,15 @@ def _expected_segment( """单个样本 pack 前的正确布局: input_ids 去掉 response 末位(EOS 只作 label)。""" assert sample.prompt_ids is not None assert sample.response_ids is not None - assert sample.response_mask is not None + assert sample.labels is not None assert sample.logprobs is not None prompt_len = len(sample.prompt_ids) input_ids = list(sample.prompt_ids) + list(sample.response_ids)[:-1] - labels = [-100] * (prompt_len - 1) + [ - resp_id if mask != 0 else -100 for resp_id, mask in zip(sample.response_ids, sample.response_mask) - ] - advantages = [0.0] * (prompt_len - 1) + [0.0 if mask == 0 else adv_val for mask in sample.response_mask] - logprobs = [0.0] * (prompt_len - 1) + list(sample.logprobs) + response_labels = list(sample.labels)[len(sample.labels) - len(sample.response_ids) :] + labels = [-100] * (prompt_len - 1) + response_labels + advantages = [0.0] * (prompt_len - 1) + [0.0 if label == -100 else adv_val for label in response_labels] + # canonical logprobs 为全序列对齐(prompt 段补 0),shift 后即 pack 布局。 + logprobs = list(sample.logprobs)[1:] return input_ids, labels, advantages, logprobs @staticmethod @@ -555,7 +871,7 @@ def _pack_and_extract(data_batches: list[dict], pack_max_length: int) -> dict: return result def _build_data_groups(self) -> list[list[RolloutState]]: - """两组样本; 组内共享 prompt(与真实 RL 组一致), 长度/mask/advantage 刻意互不相同。""" + """两组样本; 组内共享 prompt(与真实 RL 组一致), 长度/监督位/advantage 刻意互不相同。""" s1 = self._make_sample( 0, [101, 102, 103, 104], @@ -571,9 +887,9 @@ def _build_data_groups(self) -> list[list[RolloutState]]: s4 = self._make_sample(3, [301, 302, 303, 304, 305], [4001, 4002], [1, 0], [-0.5, -1.0], 8.0) return [[s1, s2], [s3, s4]] - def _prepare(self, trainer, data_groups, pack_max_length: int): - with patch("xtuner.v1.train.rl_trainer.XTUNER_DETERMINISTIC", True): - return trainer._prepare_train_data(data_groups, pack_max_length=pack_max_length) + def _convert(self, controller, data_groups, pack_max_length: int): + with patch("xtuner.v1.rl.trainer.controller.XTUNER_DETERMINISTIC", True): + return controller._convert_rollout_groups(data_groups, pack_max_length) def _assert_real_token_alignment( self, samples: list[RolloutState], packed: dict, total_len: int, packed_len: int @@ -596,12 +912,12 @@ def _assert_real_token_alignment( self.assertEqual(offset, total_len) def test_packed_fields_align_without_padding(self) -> None: - trainer = self._make_trainer() + controller = self._make_controller() data_groups = self._build_data_groups() samples = [sample for group in data_groups for sample in group] total_len = sum(len(s.prompt_ids) + len(s.response_ids) - 1 for s in samples) - data_batches, _ = self._prepare(trainer, data_groups, pack_max_length=total_len) + data_batches, _ = self._convert(controller, data_groups, pack_max_length=total_len) self.assertEqual(len(data_batches), len(samples)) for item in data_batches: @@ -611,13 +927,13 @@ def test_packed_fields_align_without_padding(self) -> None: self._assert_real_token_alignment(samples, packed, total_len, total_len) def test_packed_fields_align_with_padding(self) -> None: - trainer = self._make_trainer() + controller = self._make_controller() data_groups = self._build_data_groups() samples = [sample for group in data_groups for sample in group] total_len = sum(len(s.prompt_ids) + len(s.response_ids) - 1 for s in samples) pack_max_length = (total_len // 16 + 1) * 16 - data_batches, _ = self._prepare(trainer, data_groups, pack_max_length=pack_max_length) + data_batches, _ = self._convert(controller, data_groups, pack_max_length=pack_max_length) packed = self._pack_and_extract(data_batches, pack_max_length) self._assert_real_token_alignment(samples, packed, total_len, pack_max_length) @@ -626,12 +942,12 @@ def test_packed_fields_align_with_padding(self) -> None: self.assertTrue(all(adv == -100 for adv in packed["advantages"][total_len:])) def test_packed_fields_align_with_teacher_fields(self) -> None: - trainer = self._make_trainer(with_teacher_fields=True) + controller = self._make_controller(with_teacher_fields=True) data_groups = self._build_data_groups() samples = [sample for group in data_groups for sample in group] total_len = sum(len(s.prompt_ids) + len(s.response_ids) - 1 for s in samples) - data_batches, _ = self._prepare(trainer, data_groups, pack_max_length=total_len) + data_batches, _ = self._convert(controller, data_groups, pack_max_length=total_len) packed = self._pack_and_extract(data_batches, total_len) self._assert_real_token_alignment(samples, packed, total_len, total_len) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 364db9f09c..c4933e0f15 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -54,7 +54,6 @@ def make_rollout_state( tokens=[uid], response="" if status in (Status.ABORTED, Status.EXPIRED) else f"response {uid}", response_ids=[], - response_mask=[], finish_reason="abort" if status == Status.ABORTED else "stop", reward={"score": reward_score} if reward_score is not None else None, seq_staleness=seq_staleness, @@ -279,7 +278,6 @@ async def test_disagg_put_uses_consumer_step_for_token_expiry(self): tokens=[1, 11], response="old response", response_ids=[11], - response_mask=[1], response_model_steps=[3], logprobs=[0.1], finish_reason="stop", @@ -761,14 +759,12 @@ async def test_async_produce_strategy_rerolls_expired_state_and_preserves_fresh_ expired.response = "expired response" expired.response_ids = [11] expired.response_model_steps = [0] - expired.response_mask = None expired.reward = {"score": 0.1} fresh = make_rollout_state(901) fresh.group_id = expired.group_id fresh.response = "fresh response" fresh.response_ids = [21] fresh.response_model_steps = [5] - fresh.response_mask = None fresh.logprobs = [-0.2] fresh.reward = {"score": 0.9} diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index 737ca09d4a..691f125a97 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -50,7 +50,6 @@ def make_rollout_state( response: str | None = None, response_ids: list[int] | None = None, response_model_steps: list[int] | None = None, - response_mask: list[int] | None = None, logprobs: list[float] | None = None, reward: dict | None = None, error_msg: str | None = None, @@ -74,7 +73,6 @@ def make_rollout_state( response=response if response is not None else f"response {uid}", response_ids=response_ids, response_model_steps=list(response_model_steps) if response_model_steps is not None else None, - response_mask=list(response_mask) if response_mask is not None else [1 for _ in response_ids], logprobs=logprobs, routed_experts=routed_experts, finish_reason="stop" if status == Status.COMPLETED else None, @@ -280,7 +278,6 @@ async def test_common_put_defaults_to_retryable_expired_group(self): assert reusable.error_msg is None assert reusable.routed_experts is None assert reusable.finish_reason is None - assert reusable.response_mask is None assert reusable.mm_info is not None assert reusable.mm_info["pixel_values"] is pixel_values assert reusable.extra_fields == {"train_prompt_ids": [101, 102]} diff --git a/tests/rl/test_rl_colocate_trainer.py b/tests/rl/test_rl_colocate_trainer.py index d16446750b..502c4fb545 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -150,9 +150,6 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_ side_effect=lambda train_step, step_timer_dict: train_step % trainer._sync_weights_interval == 0 ) trainer._log_step = MagicMock() - trainer._prepare_train_data = MagicMock( - return_value=([{"seq_ctx": "fake"}], {"batch_size": 1, "rewards/mean": 1.0}) - ) trainer.rollout_controller = SimpleNamespace( shutdown_inactive_workers=SimpleNamespace( @@ -169,16 +166,19 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_ offload=MagicMock(return_value="train_offloaded"), weight_update=MagicMock(return_value="weights_updated"), fit=MagicMock( - return_value=[ - { - "rollout_is_metrics": {}, - "mismatch_metrics": {}, - "rollout_entropy": 0.0, - "train_entropy": 0.0, - "train_metrics": [], - "sft_train_metrics": {}, - } - ] + return_value=( + [ + { + "rollout_is_metrics": {}, + "mismatch_metrics": {}, + "rollout_entropy": 0.0, + "train_entropy": 0.0, + "train_metrics": [], + "sft_train_metrics": {}, + } + ], + {"batch_size": 1, "rewards/mean": 1.0}, + ) ), ) trainer.rl_health_manager = RLHealthManager( diff --git a/tests/rl/test_rl_colocate_trainer_integration.py b/tests/rl/test_rl_colocate_trainer_integration.py index 810bc220c4..8de969647f 100644 --- a/tests/rl/test_rl_colocate_trainer_integration.py +++ b/tests/rl/test_rl_colocate_trainer_integration.py @@ -27,10 +27,8 @@ TaskSpecConfig, ) from xtuner.v1.rl.evaluator import EvaluatorConfig -from xtuner.v1.data_proto.rl_data import SampleParams -from xtuner.v1.data_proto.sequence_context import SequenceContext +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status from transformers import AutoTokenizer -import torch QWEN3_PATH = os.environ["QWEN3_PATH"] ALPACA_PATH = os.environ["ALPACA_PATH"] @@ -236,7 +234,7 @@ def test_rl_train_with_sft(self): trainer = trainer_cfg.build() train_controller = trainer.train_controller - # Prepare synthetic data batches + # Prepare synthetic rollout groups tokenizer = AutoTokenizer.from_pretrained(QWEN3_PATH, trust_remote_code=True) # Create simple prompts and responses @@ -245,46 +243,46 @@ def test_rl_train_with_sft(self): ["4", "Four", "2+2=4", "The answer is 4"], ["Paris", "The capital is Paris", "Paris, France", "It's Paris"] ] + group_rewards = [1.0, 0.8, 0.9, 0.7] - data_batches = [] - for prompt, response_list in zip(prompts, responses): + rollout_groups = [] + for group_idx, (prompt, response_list) in enumerate(zip(prompts, responses)): prompt_ids = tokenizer(prompt, return_tensors='pt')['input_ids'].flatten().tolist() - rewards = torch.tensor([1.0, 0.8, 0.9, 0.7], dtype=torch.float32) - advantages = (rewards - rewards.mean()) / (rewards.std() + 1e-8) - + group = [] for i, response in enumerate(response_list): response_ids = tokenizer(response, return_tensors='pt')['input_ids'].flatten().tolist() - # Align with RLColocateTrainer._prepare_train_data(): - # - input_ids excludes last token (usually eos) of response_ids - # - shifted_labels aligns to input_ids length - input_ids = prompt_ids + response_ids[:-1] - shifted_labels = [-100] * (len(prompt_ids) - 1) + response_ids - input_ids_tensor = torch.tensor(input_ids, dtype=torch.int64).unsqueeze(0) - shifted_labels_tensor = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) - - adv_val = advantages[i].item() - # Controller._packing expects `advantage` as a list and will flatten it. - # Keep the length consistent with shifted_labels/input_ids. - advantage_list = [adv_val] * (len(prompt_ids) - 1) + [adv_val] * len(response_ids) - - data_batches.append(dict( - seq_ctx=SequenceContext.from_input_ids((input_ids_tensor,), device="cpu"), - shifted_labels=shifted_labels_tensor, - advantage=advantage_list, + # Canonical full-sequence convention (AgentLoop.canonicalize_train_fields): + # - input_ids keeps the whole prompt+response sequence + # - labels supervise the full response; TrainingController applies the + # one-position shift at conversion time + # - agent_loop_type records the producing loop's class name + group.append(RolloutState( + rollout_id=group_idx * len(response_list) + i, + group_id=group_idx, + message=[{"role": "user", "content": prompt}], + prompt_ids=prompt_ids, + response=response, + response_ids=response_ids, + agent_loop_type="SingleTurnAgentLoop", + reward={"score": group_rewards[i]}, + status=Status.COMPLETED, + input_ids=prompt_ids + response_ids, + labels=[-100] * len(prompt_ids) + response_ids, )) + rollout_groups.append(group) # RLColocateTrainer initializes by offloading train workers to CPU. # Align with RLColocateTrainer.fit() which onloads before training. train_controller.onload(target="all") # First fit and save - train_controller.fit(data_batches, pack_max_length=1024, rollout_idx=0) + train_controller.fit(rollout_groups, pack_max_length=1024, rollout_idx=0) checkpoint_path = str(work_dir / "save_test") train_controller.save(checkpoint_path, no_save_optimizer=True) # Second fit and collect metrics train_controller.onload(target="all") - log_infos = train_controller.fit(data_batches, pack_max_length=1024, rollout_idx=1) + log_infos, _ = train_controller.fit(rollout_groups, pack_max_length=1024, rollout_idx=1) efficient_attn_ratio_list = [] for log_info in log_infos: efficient_attn_ratio_list.append(log_info['sft_train_metrics']['efficient_attn_ratio']) @@ -314,7 +312,7 @@ def test_rl_train_with_sft(self): train_controller.resume(load_checkpoint_cfg) train_controller.onload(target="all") - log_infos = train_controller.fit(data_batches, pack_max_length=1024, rollout_idx=1) + log_infos, _ = train_controller.fit(rollout_groups, pack_max_length=1024, rollout_idx=1) new_efficient_attn_ratio_list = [] for log_info in log_infos: new_efficient_attn_ratio_list.append(log_info['sft_train_metrics']['efficient_attn_ratio']) diff --git a/tests/rl/test_rl_disaggregated_trainer.py b/tests/rl/test_rl_disaggregated_trainer.py index 7cff8575ad..93bdce7473 100644 --- a/tests/rl/test_rl_disaggregated_trainer.py +++ b/tests/rl/test_rl_disaggregated_trainer.py @@ -132,9 +132,6 @@ def _make_trainer(self, agent_loop_manager): trainer.eval_agent_loop_manager = SimpleNamespace(produce_batch=AsyncMock()) trainer.evaluator = MagicMock(eval_batch_size=1, run=MagicMock(return_value={"acc": 1.0})) trainer._exp_tracker = MagicMock() - trainer._prepare_train_data = MagicMock( - return_value=([{"seq_ctx": "fake"}], {"batch_size": 1, "rewards/mean": 1.0}) - ) trainer._save_trajectories = MagicMock() trainer._save_eval_trajectories = MagicMock() trainer._release_trace_sessions = MagicMock(return_value=set()) @@ -144,7 +141,12 @@ def _make_trainer(self, agent_loop_manager): trainer._maybe_save_hf = MagicMock() trainer._checkpoint_no_save_replay_buffer = False trainer.train_controller = SimpleNamespace( - fit=MagicMock(return_value=[{"train_metrics": [], "sft_train_metrics": {}}]), + fit=MagicMock( + return_value=( + [{"train_metrics": [], "sft_train_metrics": {}}], + {"batch_size": 1, "rewards/mean": 1.0}, + ) + ), onload=MagicMock(return_value="onload"), offload=MagicMock(return_value="offload"), weight_update=MagicMock(return_value="update"), diff --git a/tests/rl/test_staleness_policy.py b/tests/rl/test_staleness_policy.py index b9a746d1c8..7f70db9edf 100644 --- a/tests/rl/test_staleness_policy.py +++ b/tests/rl/test_staleness_policy.py @@ -1,6 +1,6 @@ """Staleness 配置与 mask 行为测试。 -覆盖整组过期阈值的配置和校验,以及 token 级 staleness mask 的纯函数行为。 +覆盖整组过期阈值的配置和校验,以及 token 级 staleness 往 labels 烙制的行为。 """ import unittest @@ -72,11 +72,14 @@ def test_async_strategies_precompute_token_stale_threshold(self): class TestTokenStalenessMask(unittest.TestCase): - """Token 级 staleness mask 的阈值与 semantic mask 行为。""" + """Token 级 staleness mask 的阈值与语义监督(labels 尾段)行为。""" def test_token_staleness_threshold_can_be_relaxed(self): # token threshold 放宽一个同步周期后,旧周期 token 应从 masked 变为可训练。 - for token_stale_threshold, expected in ((4, [0, 1]), (8, [1, 1])): + for token_stale_threshold, expected_mask, expected_labels in ( + (4, [0, 1], [-100, -100, -100, 4]), + (8, [1, 1], [-100, -100, 3, 4]), + ): with self.subTest(token_stale_threshold=token_stale_threshold): state = self._state(response_model_steps=[0, 4]) @@ -86,12 +89,26 @@ def test_token_staleness_threshold_can_be_relaxed(self): token_stale_threshold=token_stale_threshold, ) - self.assertEqual(masks, [expected]) - self.assertIsNone(state.response_mask) + self.assertEqual(masks, [expected_mask]) + # 过期 token 的 response 段 label 被烙成 -100;全新鲜则保持原 labels。 + self.assertEqual(state.labels, expected_labels) + + def test_repeated_calls_converge(self): + # staleness 只增不减:先在旧 step 烙一次,再在新 step 重算,最终 labels + # 与一次性按新 step 烙制的结果完全一致(replay buffer 逐轮检查同理)。 + incremental = self._state(response_model_steps=[0, 4]) + calculate_group_effective_response_masks([incremental], current_train_step=5, token_stale_threshold=4) + calculate_group_effective_response_masks([incremental], current_train_step=9, token_stale_threshold=4) + + oneshot = self._state(response_model_steps=[0, 4]) + calculate_group_effective_response_masks([oneshot], current_train_step=9, token_stale_threshold=4) - def test_token_staleness_intersects_semantic_response_mask(self): - # 最终 mask 必须同时满足 semantic mask 和 token staleness mask。 - state = self._state(response_model_steps=[0, 4], response_mask=[1, 0]) + self.assertEqual(incremental.labels, oneshot.labels) + self.assertEqual(oneshot.labels, [-100, -100, -100, -100]) + + def test_token_staleness_intersects_semantic_labels(self): + # 最终有效掩码必须同时反映语义监督(labels 尾段 -100 位)与 token staleness。 + state = self._state(response_model_steps=[0, 4], labels=[-100, -100, 3, -100]) masks = calculate_group_effective_response_masks( [state], @@ -100,11 +117,15 @@ def test_token_staleness_intersects_semantic_response_mask(self): ) self.assertEqual(masks, [[0, 0]]) + self.assertEqual(state.labels, [-100, -100, -100, -100]) - def test_rerolled_state_without_semantic_mask_uses_token_staleness_only(self): - state = reset_rollout_response(self._state(response_model_steps=[0, 4], response_mask=[0, 1])) + def test_rerolled_state_uses_token_staleness_only(self): + # 重 roll 后 canonicalize 重建训练字段:response 段全部受监督时, + # 有效掩码只由 token staleness 决定。 + state = reset_rollout_response(self._state(response_model_steps=[0, 4])) state.response_ids = [3, 4] state.response_model_steps = [4, 4] + state.labels = [-100, -100, 3, 4] masks = calculate_group_effective_response_masks( [state], @@ -112,24 +133,73 @@ def test_rerolled_state_without_semantic_mask_uses_token_staleness_only(self): token_stale_threshold=4, ) - self.assertIsNone(state.response_mask) self.assertEqual(masks, [[1, 1]]) + self.assertEqual(state.labels, [-100, -100, 3, 4]) + + def test_reset_clears_canonical_train_fields(self): + # 回归:reset 必须清掉上一轮 canonicalize 写入的 input_ids/labels, + # 否则重 roll 后基类 canonicalize 跳过重建,训练字段与新生成的 response 错位。 + state = self._state(response_model_steps=[0, 4]) + state.input_ids = [1, 2, 3, 4] + + reset_rollout_response(state) + + self.assertIsNone(state.input_ids) + self.assertIsNone(state.labels) + + def test_agentic_loop_type_is_excluded(self): + # agentic loop 类型(localhost/sandbox)产出的全序列样本不参与 token staleness,整组排除。 + agentic = self._state(response_model_steps=None, agentic=True) + agentic.input_ids = [1, 2, 3, 4] + agentic.labels = [-100, -100, 3, 4] + agentic.logprobs = [0.0, 0.0, -0.1, -0.2] + + masks = calculate_group_effective_response_masks( + [agentic], + current_train_step=5, + token_stale_threshold=4, + ) + + self.assertEqual(masks, [None]) + self.assertEqual(agentic.labels, [-100, -100, 3, 4]) + + def test_canonicalized_prompt_response_state_is_still_eligible(self): + # canonicalize 后的 prompt+response 样本(reasoning loop 产出)仍参与 staleness。 + canonical = self._state(response_model_steps=[0, 4], labels=[-100, -100, 3, -100]) + canonical.input_ids = [1, 2, 3, 4] + canonical.logprobs = [0.0, 0.0, -0.1, -0.2] + + masks = calculate_group_effective_response_masks( + [canonical], + current_train_step=5, + token_stale_threshold=4, + ) + + self.assertEqual(masks, [[0, 0]]) + self.assertEqual(canonical.labels, [-100, -100, -100, -100]) @staticmethod def _state( *, response_model_steps: list[int] | None, - response_mask: list[int] | None = None, + labels: list[int] | None = None, + agentic: bool = False, ) -> RolloutState: - return RolloutState( + state = RolloutState( rollout_id=1, group_id=1, message=[{"role": "user", "content": "prompt"}], prompt_ids=[1, 2], response_ids=[3, 4], response_model_steps=response_model_steps, - response_mask=response_mask, ) + if agentic: + state.agent_loop_type = "AgentInLocalhostLoop" + else: + # prompt+response 形态:agent_loop_type 保持 None(未记录时默认按 prompt+response 处理), + # labels 为全序列监督。 + state.labels = labels if labels is not None else [-100, -100, 3, 4] + return state if __name__ == "__main__": diff --git a/xtuner/v1/data_proto/rl_data.py b/xtuner/v1/data_proto/rl_data.py index b2e5ef43d1..cb31b03ca3 100644 --- a/xtuner/v1/data_proto/rl_data.py +++ b/xtuner/v1/data_proto/rl_data.py @@ -140,8 +140,6 @@ class RolloutState(BaseModel): teacher_targets: TeacherTargets | None = None routed_experts: np.ndarray | RayObjectRef | list[RayObjectRef] | None = None finish_reason: str | None = None - # response_mask: 记录response_ids中哪个token算loss, 与response_ids长度相同,每轮rollout在 agent_loop.generate 中覆盖写 - response_mask: list[int] | None = None # response_model_steps:记录 response_ids 中每个 token 来自哪个 model_step,与 response_ids 长度相同。 response_model_steps: list[int] | None = None # 记录该样本过期程度,即最早生成 token 的模型版本与当前训练步数的差值,数值越大表示越过期。 @@ -149,6 +147,10 @@ class RolloutState(BaseModel): input_ids: list[int] | None = None labels: list[int] | None = None + # 产出该样本的 AgentLoop 子类类名(generate_group 写入 type(self).__name__)。 + # 结合 AGENTIC_AGENT_LOOP_TYPES 可区分 agentic 全序列样本与 prompt+response(reasoning)样本; + # None(未记录,如手工构造的样本)按 prompt+response 处理。 + agent_loop_type: str | None = None # --- Judger 输出 --- reward: dict[str, Any] | None = None @@ -272,17 +274,77 @@ def reset_rollout_response(rollout_state: RolloutState) -> RolloutState: rollout_state.tokens = list(prompt_ids) if prompt_ids is not None else None rollout_state.response = "" rollout_state.response_ids = [] + # input_ids/labels 也是上一轮 canonicalize 的产物,必须一并清掉, + # 否则重生成后基类 canonicalize 见 input_ids 非空会跳过重建,留下跨轮错位的训练字段。 + rollout_state.input_ids = None + rollout_state.labels = None rollout_state.logprobs = [] rollout_state.teacher_targets = None rollout_state.routed_experts = None rollout_state.finish_reason = None - rollout_state.response_mask = None rollout_state.response_model_steps = [] rollout_state.reward = None rollout_state.error_msg = None return rollout_state +def is_valid_for_training(group_data_items: list[RolloutState], logger) -> bool: + """Checks if a group of rollout states is valid for a training step. + + Args: + group_data_items (list[RolloutState]): A list of RolloutState objects. + logger (logging.Logger): Logger used to report the invalid reason. + + Returns: + bool: True if the group is valid, False otherwise. + + NOTE: Why this check is needed: + - For system fault tolerance, this check is performed at rollout / dataflow + time, but we still do it here to ensure training data integrity. + - 'filtered'/'failed': These items are fundamentally broken or incomplete and + should not be used for training. + - 'aborted': These items represent rollouts that were stopped + prematurely. Using such partial data could lead the model to learn + undesirable behaviors (e.g., stopping generation too early). + - Empty response/response_ids: The model's generated response is the core + of the training data for RL algorithms like PPO. If the response is + missing, there is nothing to compute rewards on or to train the model with. + """ + is_abort = any(item.status == Status.ABORTED for item in group_data_items) + is_filtered = any(item.status == Status.FILTERED for item in group_data_items) + is_failed = any(item.status == Status.FAILED for item in group_data_items) + if is_filtered or is_failed or is_abort: + logger.warning( + f"Invalid dataflow group found during training, rollout state filtered: {is_filtered}, failed: {is_failed}, aborted: {is_abort}." + ) + return False + for item in group_data_items: + if item.input_ids is not None: + input_ids_valid = len(item.input_ids) > 1 + labels_valid = item.labels is not None and len(item.labels) == len(item.input_ids) + logprobs_valid = item.logprobs is None or len(item.logprobs) == len(item.input_ids) + if not input_ids_valid or not labels_valid or not logprobs_valid: + logger.warning( + "Invalid dataflow item found during training: input_ids, labels, and logprobs lengths mismatch." + ) + return False + continue + + response_valid = item.response is not None and len(item.response) > 0 + ids_valid = item.response_ids is not None and len(item.response_ids) > 0 + if not ids_valid: + # NOTE: `response_ids` is the critical field for token-in-token-out mode, so we ensure it's not empty. + logger.warning( + "Invalid dataflow item found during training: no response or response_ids and skip this item." + ) + return False + if not response_valid: + # NOTE: check valid response string for judger inputs + logger.warning("Invalid dataflow item found during training: empty response string and skip this item.") + return False + return True + + def get_group_status(rollout_states: list[RolloutState]) -> Status: """Get the group status based on the individual rollout states. @@ -351,6 +413,11 @@ def refresh_seq_staleness(group: list[RolloutState], current_train_step: int) -> return group +# 产自这些 AgentLoop 子类(类名记录在 RolloutState.agent_loop_type)的样本是 agentic 全序列样本: +# labels 无 prompt 段偏移语义,不参与 token 级 staleness。自定义 agentic loop 需把类名加入此集合。 +AGENTIC_AGENT_LOOP_TYPES: frozenset[str] = frozenset({"AgentInLocalhostLoop", "AgentInSandboxLoop"}) + + def _calculate_effective_response_mask( rollout_state: RolloutState, *, @@ -359,6 +426,9 @@ def _calculate_effective_response_mask( ) -> list[int]: """Calculate the response mask after applying token staleness. + The semantic mask is recovered from the labels tail (response segment): supervised + tokens carry their target id, masked-out tokens are ``-100``. + Args: rollout_state (RolloutState): Rollout sample whose response token provenance is evaluated. current_train_step (int): Trainer step that will consume the sample. @@ -367,13 +437,14 @@ def _calculate_effective_response_mask( Returns: list[int]: The semantic response mask intersected with the token-staleness mask. """ + labels = cast(list[int], rollout_state.labels) response_ids = cast(list[int], rollout_state.response_ids) response_model_steps = cast(list[int], rollout_state.response_model_steps) - # semantic mask: 在 agent_loop 中根据是否是 LLM 产生的 token 来 mask 的结果 - semantic_mask = rollout_state.response_mask - if semantic_mask is None: - semantic_mask = [1] * len(response_ids) + # semantic mask: 从 labels 的 response 段恢复(-100 即语义上不参与 loss 的 token)。 + # response 段偏移 = len(labels) - len(response_ids),即 prompt 段长度。 + offset = len(labels) - len(response_ids) + semantic_mask = [int(label != -100) for label in labels[offset:]] # token_staleness_mask: 根据 token 的新鲜程度来 mask token_staleness_mask = [ @@ -393,25 +464,46 @@ def calculate_group_effective_response_masks( current_train_step: int, token_stale_threshold: int | None, ) -> list[list[int] | None]: - """Calculate token-staleness masks for the applicable states in a group. + """Calculate a group's effective masks and bake token staleness into its + labels. + + For each eligible state, response labels whose token staleness reaches the threshold + are cleared to ``-100`` in place. Clearing only ever extends: staleness grows + monotonically with the trainer step, so repeated calls (e.g. replay-buffer expiry + checks followed by the train-batch bake) converge to the same labels. Each returned + mask is the semantic response mask intersected with the token-staleness mask. + ``None`` means token staleness is disabled or does not apply to that state. Agentic + full-sequence states (``agent_loop_type`` in ``AGENTIC_AGENT_LOOP_TYPES``) receive + ``None``. + + Args: + group (list[RolloutState]): Rollout group updated in place. + current_train_step (int): Trainer step that will consume the group. + token_stale_threshold (int | None): Maximum token staleness, measured in trainer + steps, allowed for training. ``None`` disables token staleness. - Each ``None`` means that token staleness is disabled or does not apply to - that state. Agentic groups currently return one ``None`` per state. + Returns: + list[list[int] | None]: Per-state effective masks. """ if token_stale_threshold is None: return [None] * len(group) - if any(item.input_ids is not None or item.labels is not None for item in group): + if any(item.agent_loop_type in AGENTIC_AGENT_LOOP_TYPES for item in group): return [None] * len(group) - return [ - ( - None - if not item.response_ids or (item.response_mask is not None and not any(item.response_mask)) - else _calculate_effective_response_mask( - item, - current_train_step=current_train_step, - token_stale_threshold=token_stale_threshold, - ) + masks: list[list[int] | None] = [] + for item in group: + if item.labels is None or not item.response_ids or len(item.labels) <= len(item.response_ids): + masks.append(None) + continue + effective_mask = _calculate_effective_response_mask( + item, + current_train_step=current_train_step, + token_stale_threshold=token_stale_threshold, ) - for item in group - ] + # prompt 段长度由 labels 与 response 段(有效掩码)长度差推导。 + offset = len(item.labels) - len(effective_mask) + for i, mask_value in enumerate(effective_mask): + if mask_value == 0: + item.labels[offset + i] = -100 + masks.append(effective_mask) + return masks diff --git a/xtuner/v1/rl/agent_loop/agent_loop.py b/xtuner/v1/rl/agent_loop/agent_loop.py index 823942d05a..a25cc579de 100644 --- a/xtuner/v1/rl/agent_loop/agent_loop.py +++ b/xtuner/v1/rl/agent_loop/agent_loop.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, TypeAlias, cast, overload import ray +import torch from pydantic import BaseModel, ConfigDict from ray.actor import ActorClass, ActorProxy from ray.util.placement_group import PlacementGroup @@ -274,13 +275,80 @@ async def generate_one(state: RolloutState) -> RolloutState: state = await self._teacher_scorer.on_sample_ready(state) return state + for state in rollout_state: + state.agent_loop_type = type(self).__name__ group = list(await asyncio.gather(*(create_task(generate_one(state)) for state in rollout_state))) if self.judger is not None and self.enable_batch_judge: if all(sample.status == Status.COMPLETED for sample in group): group = await self.run_judger(group) group = maybe_filter_invalid_sample(group, self.is_valid_sample_fn, self.logger) - return await self._teacher_scorer.on_group_ready(group) + group = await self._teacher_scorer.on_group_ready(group) + return self.canonicalize_train_fields(group) + + def canonicalize_train_fields(self, group: list[RolloutState]) -> list[RolloutState]: + """Canonicalize the training token fields of one rollout group. + + Prompt+response samples without ``input_ids`` are assembled in place into the + unified full-sequence convention: ``input_ids``/``labels``/``logprobs`` share one + length, where ``labels[j]`` is the prediction target of the token at position + ``j`` and the train controller applies the one-position shift at conversion time. + Samples that already carry ``input_ids`` satisfy the convention and are left + untouched. Non-completed samples are skipped. Failures mark the single sample as + ``Status.FAILED`` instead of raising into the producer. Loops that assemble their + own full sequences (e.g. the agentic localhost/sandbox loops) override this + method with their own loop-specific implementation instead of relying on this + default. + + Args: + group (list[RolloutState]): One group of rollout states after generation, + judging, and filtering. + + Returns: + list[RolloutState]: The same group with canonical training fields. + """ + for state in group: + if state.status != Status.COMPLETED or state.input_ids is not None: + continue + try: + is_vlm_model = "train_prompt_ids" in state.extra_fields + if is_vlm_model: + # TODO(hha): VLM, 不好的设计,后续要去掉 + prompt_ids = state.extra_fields["train_prompt_ids"] + else: + prompt_ids = state.prompt_ids + assert prompt_ids is not None and len(prompt_ids) > 0, ( + f"Prompt ids cannot be None or empty in data: {state}" + ) + + response_ids: list[int] + if state.response_ids is not None: + resp_ids_raw = state.response_ids + if isinstance(resp_ids_raw, torch.Tensor): + response_ids = resp_ids_raw.flatten().tolist() + else: + response_ids = cast(list[int], resp_ids_raw) + else: + assert state.response is not None, "response item cannot be None" + response_ids = self.tokenizer(state.response, return_tensors="pt")["input_ids"].flatten().tolist() + # 归一化写回:Tensor 展平为 list、tokenizer 兜底结果也落到 response_ids,后续消费者零分支。 + state.response_ids = response_ids + + if state.logprobs is not None: + assert len(state.logprobs) == len(response_ids), ( + f"{len(state.logprobs)} vs {len(response_ids)}, data: {state}" + ) + # 只有 response 部分有 logprobs, 前面补 0 对齐全序列 + state.logprobs = [0.0] * len(prompt_ids) + list(state.logprobs) + + # 返回的 routed_experts 不包括 eos 的值,实际上也不需要,训练输入会去掉最后一位 + state.input_ids = list(prompt_ids) + response_ids + state.labels = [-100] * len(prompt_ids) + list(response_ids) + except Exception as exc: + self.logger.error(f"Canonicalize train fields failed for rollout_id={state.rollout_id}: {exc}") + state.status = Status.FAILED + state.error_msg = f"canonicalize_train_fields failed: {exc}" + return group @overload async def run_judger(self, rollout_state: RolloutState) -> RolloutState: ... diff --git a/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py b/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py index 13044d4f35..4108a48765 100644 --- a/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py +++ b/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py @@ -152,12 +152,19 @@ async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> Rollou final_response_mask = final_response_mask[:max_len] final_logprobs = final_logprobs[:max_len] + prompt_ids = rollout_state.prompt_ids + assert prompt_ids is not None and len(prompt_ids) > 0, ( + f"Prompt ids cannot be None or empty in data: {rollout_state}" + ) rollout_state.response_ids = final_response_ids - rollout_state.response_mask = final_response_mask rollout_state.logprobs = final_logprobs + rollout_state.input_ids = list(prompt_ids) + final_response_ids + rollout_state.labels = [-100] * len(prompt_ids) + [ + resp_id if mask_value else -100 for resp_id, mask_value in zip(final_response_ids, final_response_mask) + ] rollout_state.response = self.tokenizer.decode(rollout_state.response_ids) - assert len(rollout_state.response_ids) == len(rollout_state.response_mask) == len(rollout_state.logprobs), ( - f"{len(rollout_state.response_ids)} vs {len(rollout_state.response_mask)} vs {len(rollout_state.logprobs)}" + assert len(rollout_state.input_ids) == len(rollout_state.labels) == len(rollout_state.logprobs), ( + f"{len(rollout_state.input_ids)} vs {len(rollout_state.labels)} vs {len(rollout_state.logprobs)}" ) if self.judger is not None and not self.enable_batch_judge: rollout_state = await self.run_judger(rollout_state) diff --git a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py index bdba686585..a4ac2da918 100644 --- a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py +++ b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py @@ -145,6 +145,7 @@ async def generate_one(state: RolloutState) -> RolloutState: tasks: list[asyncio.Task[RolloutState]] = [] for state in rollout_state: + state.agent_loop_type = type(self).__name__ state.sample_params = self.sample_params task = create_task(generate_one(state)) tasks.append(task) @@ -152,7 +153,42 @@ async def generate_one(state: RolloutState) -> RolloutState: samples = await asyncio.gather(*tasks) samples = _drop_failed_train_samples(samples, self.mode) samples = maybe_filter_invalid_sample(samples, self.is_valid_sample_fn, self.logger) - return await self._teacher_scorer.on_group_ready(samples) + samples = await self._teacher_scorer.on_group_ready(samples) + return self.canonicalize_train_fields(samples) + + def canonicalize_train_fields(self, group: list[RolloutState]) -> list[RolloutState]: + """Validate the canonical training fields of one localhost agentic + rollout group. + + This loop fills ``input_ids``/``labels``/``logprobs`` from the trace store as one + unshifted full sequence, so there is nothing left to assemble. Canonicalization + only re-checks the unified convention (``labels`` present and of equal length, + ``logprobs`` either ``None`` or of equal length) and marks a violating sample + ``Status.FAILED`` instead of raising into the producer. States without + ``input_ids`` (eval-mode states) carry no training fields and are skipped. + + Args: + group (list[RolloutState]): One group of rollout states after generation, + judging, and filtering. + + Returns: + list[RolloutState]: The same group with validated training fields. + """ + for state in group: + if state.status != Status.COMPLETED or state.input_ids is None: + continue + try: + labels = state.labels + if labels is None or len(labels) != len(state.input_ids): + expected = 0 if labels is None else len(labels) + raise ValueError(f"labels length mismatch: {expected} vs {len(state.input_ids)}") + if state.logprobs is not None and len(state.logprobs) != len(state.input_ids): + raise ValueError(f"logprobs length mismatch: {len(state.logprobs)} vs {len(state.input_ids)}") + except Exception as exc: + self.logger.error(f"Canonicalize train fields failed for rollout_id={state.rollout_id}: {exc}") + state.status = Status.FAILED + state.error_msg = f"canonicalize_train_fields failed: {exc}" + return group async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: try: @@ -278,7 +314,6 @@ def _fill_eval_rollout_state(self, rollout_state: RolloutState, item: AgentRollo rollout_state.response_ids = None rollout_state.logprobs = None rollout_state.routed_experts = None - rollout_state.response_mask = None rollout_state.response_model_steps = None rollout_state.extra_fields["agent_status"] = item.status.value if item.error is not None: diff --git a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py index 14dec5bf74..f40e2ea8af 100644 --- a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py +++ b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py @@ -240,6 +240,7 @@ async def generate_one(state: RolloutState) -> list[RolloutState]: pending_tasks = [] for state in rollout_state: + state.agent_loop_type = type(self).__name__ state.sample_params = self.sample_params task = create_task(generate_one(state)) pending_tasks.append(task) @@ -248,7 +249,43 @@ async def generate_one(state: RolloutState) -> list[RolloutState]: samples = [sample for sample_group in sample_groups for sample in sample_group] samples = _drop_failed_train_samples(samples, self.mode) samples = maybe_filter_invalid_sample(samples, self.is_valid_sample_fn, self.logger) - return await self._teacher_scorer.on_group_ready(samples) + samples = await self._teacher_scorer.on_group_ready(samples) + return self.canonicalize_train_fields(samples) + + def canonicalize_train_fields(self, group: list[RolloutState]) -> list[RolloutState]: + """Validate the canonical training fields of one sandbox agentic + rollout group. + + Each segment state carries ``input_ids``/``labels``/``logprobs`` exported from + the trace store as one unshifted full sequence, so there is nothing left to + assemble. Canonicalization only re-checks the unified convention (``labels`` + present and of equal length, ``logprobs`` either ``None`` or of equal length) + per segment state and marks a violating sample ``Status.FAILED`` instead of + raising into the producer. States without ``input_ids`` (eval-mode states) carry + no training fields and are skipped. + + Args: + group (list[RolloutState]): One flattened group of segment rollout states + after generation, judging, and filtering. + + Returns: + list[RolloutState]: The same group with validated training fields. + """ + for state in group: + if state.status != Status.COMPLETED or state.input_ids is None: + continue + try: + labels = state.labels + if labels is None or len(labels) != len(state.input_ids): + expected = 0 if labels is None else len(labels) + raise ValueError(f"labels length mismatch: {expected} vs {len(state.input_ids)}") + if state.logprobs is not None and len(state.logprobs) != len(state.input_ids): + raise ValueError(f"logprobs length mismatch: {len(state.logprobs)} vs {len(state.input_ids)}") + except Exception as exc: + self.logger.error(f"Canonicalize train fields failed for rollout_id={state.rollout_id}: {exc}") + state.status = Status.FAILED + state.error_msg = f"canonicalize_train_fields failed: {exc}" + return group # NOTE: A single sandbox session may yield multiple trainable segments, so this returns a list # rather than the base class's single RolloutState. The base contract is never exercised for @@ -370,7 +407,6 @@ def _fill_eval_rollout_state(self, rollout_state: RolloutState, item: AgentRollo rollout_state.response_ids = None rollout_state.logprobs = None rollout_state.routed_experts = None - rollout_state.response_mask = None rollout_state.response_model_steps = None rollout_state.extra_fields["origin_data_source"] = item.data_source rollout_state.extra_fields["agent_status"] = item.status.value diff --git a/xtuner/v1/rl/agent_loop_manager/disagg_producer.py b/xtuner/v1/rl/agent_loop_manager/disagg_producer.py index 76503ebe17..59a6762612 100644 --- a/xtuner/v1/rl/agent_loop_manager/disagg_producer.py +++ b/xtuner/v1/rl/agent_loop_manager/disagg_producer.py @@ -259,9 +259,9 @@ class DisaggAsyncProduceStrategyConfig(DisaggProduceStrategyConfig): response token may lag behind before it is masked out of the loss. ``None`` disables token-level masking, ``0`` accepts only tokens produced within the current sync period, and ``N`` allows ``N`` - extra periods. Partially stale responses have their - ``response_mask`` reduced before training. If a state has no - trainable response token left, the state expires and its group + extra periods. Partially stale responses have the corresponding + response labels cleared to ``-100`` before training. If a state has + no trainable response token left, the state expires and its group enters the expired-group lifecycle. Defaults to None. tail_batch_trigger_size (int): Expired-group rerollout policy. ``-1`` disables rerollout and terminally discards expired groups, ``0`` diff --git a/xtuner/v1/rl/agent_loop_manager/produce_utils.py b/xtuner/v1/rl/agent_loop_manager/produce_utils.py index 1a2f8e365f..8214bc581a 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -593,14 +593,11 @@ async def take_train_batch( if task.token_stale_threshold is None: continue for group in batch_by_task.get(task.task_name, []): - effective_masks = calculate_group_effective_response_masks( + calculate_group_effective_response_masks( group, current_train_step=current_train_step, token_stale_threshold=task.token_stale_threshold, ) - for rollout_state, effective_mask in zip(group, effective_masks): - if effective_mask is not None: - rollout_state.response_mask = effective_mask if hasattr(progress, "mark_consumed"): progress.mark_consumed(consumed_counts) diff --git a/xtuner/v1/rl/agent_loop_manager/producer.py b/xtuner/v1/rl/agent_loop_manager/producer.py index 2f7f660e70..7f6186e116 100644 --- a/xtuner/v1/rl/agent_loop_manager/producer.py +++ b/xtuner/v1/rl/agent_loop_manager/producer.py @@ -196,9 +196,9 @@ class AsyncProduceStrategyConfig(ProduceStrategyConfig): response token may lag behind before it is masked out of the loss. ``None`` disables token-level masking, ``0`` accepts only tokens produced within the current sync period, and ``N`` allows ``N`` - extra periods. Partially stale responses have their - ``response_mask`` reduced before training. If a state has no - trainable response token left, the state expires and its group + extra periods. Partially stale responses have the corresponding + response labels cleared to ``-100`` before training. If a state has + no trainable response token left, the state expires and its group enters the expired-group lifecycle. Defaults to None. tail_batch_trigger_size (int): Expired-group rerollout policy. ``-1`` disables rerollout and terminally discards expired groups, ``0`` diff --git a/xtuner/v1/rl/replay_buffer.py b/xtuner/v1/rl/replay_buffer.py index 61a2abacce..4f033ddbb9 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -459,7 +459,7 @@ def _apply_staleness_lifecycle( # 1. update seq-level staleness refresh_seq_staleness(group, current_train_step) - # 2. calculate token-level staleness + # 2. bake token-level staleness into labels and use the masks for expiry decisions token_level_effective_masks = calculate_group_effective_response_masks( group, current_train_step=current_train_step, diff --git a/xtuner/v1/rl/trainer/controller.py b/xtuner/v1/rl/trainer/controller.py index 17f2832067..14813351dc 100644 --- a/xtuner/v1/rl/trainer/controller.py +++ b/xtuner/v1/rl/trainer/controller.py @@ -1,27 +1,36 @@ +from __future__ import annotations + import math import os -from typing import Any, Literal, TypedDict +import random +from typing import TYPE_CHECKING, Any, Literal, TypedDict, cast +import numpy as np import ray import torch from typing_extensions import NotRequired +from xtuner.v1.data_proto.rl_data import AGENTIC_AGENT_LOOP_TYPES, RolloutState, is_valid_for_training from xtuner.v1.data_proto.sequence_context import SequenceContext from xtuner.v1.model.compose.base import BaseComposeConfig +from xtuner.v1.rl.distillation import DistillationTrainerAdapter from xtuner.v1.rl.utils import free_object_refs from xtuner.v1.train.trainer import LoadCheckpointConfig -from xtuner.v1.utils import get_logger +from xtuner.v1.utils import XTUNER_DETERMINISTIC, get_logger from .worker import TrainingWorker, WorkerLogItem +if TYPE_CHECKING: + from xtuner.v1.rl.advantage.base import AdvantageEstimator + TRAIN_RAY_GET_TIMEOUT = os.getenv("XTUNER_TRAIN_RAY_GET_TIMEOUT", 5 * 3600) # default 5 hours class ColateItem(TypedDict): seq_ctx: SequenceContext shifted_labels: torch.Tensor - advantage: float + advantage: list[float] rollout_logprobs: torch.Tensor | None teacher_logprobs: NotRequired[torch.Tensor | None] target_token_ids: NotRequired[torch.Tensor | None] @@ -41,6 +50,62 @@ class PackedBatch(TypedDict): teacher_indices: torch.Tensor | None +def get_train_seq_ctx( + input_ids: torch.LongTensor, + position_ids: np.ndarray | None = None, + multimodal_train_info: dict | None = None, + len_response_ids: int = 0, +) -> SequenceContext: + """Build a CPU ``SequenceContext`` for one training sample. + + Args: + input_ids (torch.LongTensor): Model input tokens with shape ``(1, seq_len)``. + position_ids (np.ndarray | None): Optional position ids. A 3D array triggers + the VLM MRoPE layout and the response segment is appended. + multimodal_train_info (dict | None): Optional multimodal payload with + ``pixel_values``, ``image_grid_thw`` and ``num_img_tokens``. + len_response_ids (int): Response segment length used to extend 3D position ids. + + Returns: + SequenceContext: The CPU sequence context of the sample. + """ + seq_ctx = SequenceContext.from_input_ids((input_ids,), device="cpu") + position_ids = _to_cpu_tensor(position_ids, dtype=torch.long) + if position_ids is not None and len(position_ids.shape) == 3: + # Match get_rope_index_3: response text continues from a single global + # max(T, H, W), not per-axis maxima. Per-axis max diverges when the + # prompt ends on image tokens (T≈0 while H/W are large). + max_value = position_ids.amax() + response_position_ids = ( + ( + torch.arange( + 1, + len_response_ids + 1, + device=position_ids.device, + dtype=position_ids.dtype, + ) + + max_value + ) + .view(1, 1, -1) + .expand(3, 1, -1) + ) + position_ids = torch.cat([position_ids, response_position_ids], dim=-1) + seq_ctx.position_ids = position_ids # type: ignore[assignment] + assert position_ids.size(-1) == input_ids.size(-1) + + if multimodal_train_info: + seq_ctx.pixel_values = multimodal_train_info.get("pixel_values") + seq_ctx.image_grid_thw = _to_cpu_tensor(multimodal_train_info.get("image_grid_thw"), dtype=torch.long) + num_img_tokens = multimodal_train_info.get("num_img_tokens") + if num_img_tokens is not None: + seq_ctx.num_img_tokens = [num_img_tokens] + return seq_ctx + + +def _stat_tensor(values: list[float]) -> torch.Tensor: + return torch.tensor(values).float() if values else torch.tensor([0.0]).float() + + def _summarize_process_group_results(results: list[dict[str, Any]]) -> str: if not results: return "ranks=0" @@ -84,8 +149,16 @@ def _verify_packed_alignment(packed_batch: PackedBatch) -> None: class TrainingController: - def __init__(self, workers: list[TrainingWorker]) -> None: + def __init__( + self, + workers: list[TrainingWorker], + advantage_estimator: AdvantageEstimator | None = None, + distillation: DistillationTrainerAdapter | None = None, + ) -> None: self.workers = workers + self.advantage_estimator = advantage_estimator + self.distillation = distillation if distillation is not None else DistillationTrainerAdapter(None) + self.task_adv_weight = self.distillation.task_adv_weight self.logger = get_logger() # TODO(hha): 这个逻辑不够通用,应该复用 sft 函数,从而支持 expand soft pack @@ -273,7 +346,219 @@ def _grouped_by_max_length(self, packed_data_batches): # 排序后这条 pack 会被放在最前面,导致 rank0 的第一个 step 消耗的有效 token 数往往少于其他 rank,是正常现象。 return sorted(packed_data_batches, key=lambda x: x["seq_ctx"].max_length_q, reverse=True) - def fit(self, data_batches: list[ColateItem], pack_max_length: int, rollout_idx: int) -> list[WorkerLogItem]: + def _convert_rollout_groups( + self, + rollout_groups: list[list[RolloutState]], + pack_max_length: int, + raw_rewards_sum: float = 0.0, + raw_rewards_count: int = 0, + ) -> tuple[list[ColateItem], dict[str, float]]: + """Convert rollout groups into packable training items and step + statistics. + + Samples arrive canonicalized from the agent loop (``input_ids``/``labels``/ + ``logprobs`` share one full-sequence length) with token staleness already baked + into ``labels``. This method shifts by one position, builds tensors and sequence + contexts, computes group-level advantages from session clustered rewards, and + summarizes the ``data_info`` metrics. + + Args: + rollout_groups (list[list[RolloutState]]): Rollout groups selected for training. + pack_max_length (int): Maximum token length per sample accepted for packing. + raw_rewards_sum (float): Producer-side raw reward sum used by ``raw_rewards/mean``. + raw_rewards_count (int): Producer-side raw reward count used by ``raw_rewards/mean``. + + Returns: + tuple[list[ColateItem], dict[str, float]]: Packable training items and the step + statistics dictionary. + """ + has_3d_position = any( + state.position_ids is not None and state.position_ids.ndim == 3 + for group in rollout_groups + for state in group + ) + + # Per-session rewards for distribution metrics. Agentic sessions may split into several + # trainable segments that share one reward; counting that reward once per session keeps + # rewards/* from being weighted by segment count. + cluster_rewards_list: list[float] = [] + distillation_reward_observations: list[tuple[RolloutState, float]] = [] + advantages_list: list[float] = [] + prompt_len_list: list[float] = [] + response_len_list: list[float] = [] + tool_turns_list: list[int] = [] + training_tokens = 0 + training_samples = 0 + + data_batches: list[ColateItem] = [] + + for group in rollout_groups: + if not is_valid_for_training(group, self.logger): + self.logger.error(f"Skip one data group {group} due to rollout failed or empty response.") + continue + training_samples += len(group) + + # Collect rewards independently from task-advantage computation. Pure OPD may omit + # rewards entirely; when rewards are present they remain useful observability signals. + # session_id 只由 agentic loop / XTUNER_DETERMINISTIC 写入;普通 RL 回退到 rollout_id。 + rewards_by_session: dict[Any, float] = {} + session_representatives: list[RolloutState] = [] + cluster_indices: list[int | None] = [] + for state in group: + # 有可能有重复,但是没有其他更好办法 + turns = state.extra_fields.get("agent_tool_turns") + if isinstance(turns, int): + tool_turns_list.append(turns) + if state.reward is None or "score" not in state.reward: + if self.task_adv_weight > 0: + raise ValueError( + f"Reward is missing or does not contain 'score' key in data: {state}, " + f"but task_adv_weight={self.task_adv_weight} > 0" + ) + cluster_indices.append(None) + continue + reward = float(state.reward["score"]) + session_key = state.session_id if state.session_id is not None else state.rollout_id + if session_key not in rewards_by_session: + rewards_by_session[session_key] = reward + session_representatives.append(state) + cluster_indices.append(len(session_representatives) - 1) + + cluster_rewards = list(rewards_by_session.values()) + cluster_rewards_list.extend(cluster_rewards) + distillation_reward_observations.extend(zip(session_representatives, cluster_rewards)) + + if self.task_adv_weight != 0 and self.advantage_estimator is not None: + cluster_advantages = self.advantage_estimator.compute( + torch.tensor(cluster_rewards, dtype=torch.float32), session_representatives + ) + sample_advantages = [ + float(cluster_advantages[index].item()) if index is not None else 0.0 for index in cluster_indices + ] + else: + sample_advantages = [0.0] * len(group) + + for i, state in enumerate(group): + input_ids = state.input_ids + labels = state.labels + assert input_ids is not None and labels is not None and len(input_ids) == len(labels), ( + f"Rollout state is not canonicalized for training: {state}" + ) + input_len = len(input_ids) + shifted_labels = labels[1:] + + input_ids_t = cast(torch.LongTensor, torch.tensor(input_ids[:-1], dtype=torch.int64).unsqueeze(0)) + shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) + + rollout_logprobs: torch.Tensor | None = None + if state.logprobs is not None: + assert len(state.logprobs) == input_len, f"{len(state.logprobs)} vs {input_len}, data: {state}" + rollout_logprobs = torch.tensor(state.logprobs[1:], dtype=torch.float32).unsqueeze(0) + assert rollout_logprobs.size() == shifted_labels_t.size(), ( + f"{rollout_logprobs.size()} vs {shifted_labels_t.size()}" + ) + + # Keep the advantage layout aligned with input_ids (response excludes EOS). + # Prompt positions predict prompt tokens whose shifted labels are -100, so their + # advantage is 0; the last entry of response_ids (EOS) is only a label, never an input. + advantage_val = sample_advantages[i] + actual_advantages = [0.0 if label == -100 else advantage_val for label in shifted_labels] + advantages_list.extend(advantage_val for label in shifted_labels if label != -100) + + assert input_len - 1 <= pack_max_length, f"{input_len - 1} vs {pack_max_length}" + training_tokens += input_len - 1 + + # prompt+response(reasoning)样本按 prompt/response 原始长度统计分母; + # agentic 样本按 -100 占位统计。 + if state.agent_loop_type in AGENTIC_AGENT_LOOP_TYPES or state.response_ids is None: + prompt_len = sum(label == -100 for label in shifted_labels) + response_len = len(shifted_labels) - prompt_len + else: + response_len = len(state.response_ids) + prompt_len = input_len - response_len + prompt_len_list.append(prompt_len) + response_len_list.append(response_len) + + position_ids = state.position_ids + if has_3d_position and (position_ids is None or position_ids.ndim != 3): + seq_len = input_ids_t.size(-1) + text_position_ids = np.arange(seq_len, dtype=np.int64).reshape(1, 1, -1) + # Mixed batches use three-axis MRoPE position IDs. Repeat the normal text + # positions on every axis so text samples have the same rank as the VLM + # samples when their sequence contexts are packed together. + position_ids = np.broadcast_to(text_position_ids, (3, 1, seq_len)).copy() + multimodal_train_info = cast(dict | None, state.mm_info) + len_response_ids = ( + len(state.response_ids) - 1 + if state.response_ids is not None and state.agent_loop_type not in AGENTIC_AGENT_LOOP_TYPES + else 0 + ) + seq_ctx = get_train_seq_ctx(input_ids_t, position_ids, multimodal_train_info, len_response_ids) + seq_ctx.rollout_routed_experts = state.routed_experts # type: ignore[assignment] # n,layer*expert + + teacher_fields = self.distillation.rollout_teacher_targets(state, shifted_labels=shifted_labels) + + data_dict: ColateItem = { + "seq_ctx": seq_ctx, + "shifted_labels": shifted_labels_t, + "advantage": actual_advantages, + "rollout_logprobs": rollout_logprobs, + } + data_dict.update(cast(ColateItem, teacher_fields)) + data_batches.append(data_dict) + + if not XTUNER_DETERMINISTIC: + random.shuffle(data_batches) + + # rewards/* report the per-session reward distribution; batch_size/training_samples + # count the valid rollout segments included in training. + rewards_t = _stat_tensor(cluster_rewards_list) + advantages_t = _stat_tensor(advantages_list) + prompt_len_t = _stat_tensor(prompt_len_list) + response_len_t = _stat_tensor(response_len_list) + + raw_rewards_mean = raw_rewards_sum / raw_rewards_count if raw_rewards_count > 0 else rewards_t.mean().item() + info_dict: dict[str, float] = { + "batch_size": training_samples, + "training_samples": training_samples, + "training_tokens": training_tokens, + "rewards/mean": rewards_t.mean().item(), + "rewards/min": rewards_t.min().item(), + "rewards/max": rewards_t.max().item(), + "raw_rewards/mean": raw_rewards_mean, + "advantages/mean": advantages_t.mean().item(), + "advantages/min": advantages_t.min().item(), + "advantages/max": advantages_t.max().item(), + "response_len/mean": response_len_t.mean().item(), + "response_len/min": response_len_t.min().item(), + "response_len/max": response_len_t.max().item(), + "response_len/std": response_len_t.std().item(), + "prompt_len/mean": prompt_len_t.mean().item(), + "prompt_len/min": prompt_len_t.min().item(), + "prompt_len/max": prompt_len_t.max().item(), + } + info_dict.update(self.distillation.reward_scalars(distillation_reward_observations)) + if tool_turns_list: + tool_turns_t = torch.tensor(tool_turns_list, dtype=torch.float32) + info_dict["tool_turns/mean"] = tool_turns_t.mean().item() + info_dict["tool_turns/min"] = float(tool_turns_t.min().item()) + info_dict["tool_turns/max"] = float(tool_turns_t.max().item()) + return data_batches, info_dict + + def fit( + self, + rollout_groups: list[list[RolloutState]], + pack_max_length: int, + rollout_idx: int, + raw_rewards_sum: float = 0.0, + raw_rewards_count: int = 0, + ) -> tuple[list[WorkerLogItem], dict[str, float]]: + data_batches, data_info = self._convert_rollout_groups( + rollout_groups, + pack_max_length, + raw_rewards_sum=raw_rewards_sum, + raw_rewards_count=raw_rewards_count, + ) has_rollout_routed_experts = False language_cfg = None if data_batches[0]["seq_ctx"].rollout_routed_experts is not None: @@ -405,7 +690,7 @@ def fit(self, data_batches: list[ColateItem], pack_max_length: int, rollout_idx: free_object_refs(free_pixel_value_refs) del data_batch_refs del packed_data_batches - return log_infos + return log_infos, data_info def offload(self, target: Literal["model", "optimizer", "all"] = "all"): if target == "model": @@ -490,3 +775,10 @@ def save(self, dcp_dir: str, no_save_optimizer: bool = False): handles = [worker.save.remote(dcp_dir, no_save_optimizer) for worker in self.workers] # type: ignore ray.get(handles, timeout=TRAIN_RAY_GET_TIMEOUT) return + + +def _to_cpu_tensor(value: np.ndarray | None, *, dtype: torch.dtype | None = None) -> torch.Tensor | None: + if value is None: + return None + assert isinstance(value, np.ndarray), f"Expected np.ndarray, got {type(value)}" + return torch.as_tensor(value, dtype=dtype, device="cpu") diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index 0193dd56ea..0b6df62b3c 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import contextlib import json import math @@ -8,7 +10,6 @@ from typing import ( TYPE_CHECKING, Iterable, - List, Sequence, TypedDict, cast, @@ -18,6 +19,8 @@ if TYPE_CHECKING: from ray.util.placement_group import PlacementGroup + from xtuner.v1.rl.advantage.base import AdvantageEstimator + import numpy as np import ray import torch @@ -43,6 +46,7 @@ from xtuner.v1.model.utils.misc import ModelForwardExtraLogInfo from xtuner.v1.profiler import profiling_memory, profiling_time from xtuner.v1.rl.distillation import ( + DistillationTrainerAdapter, TrainTeacherManager, TrainTeacherManagerConfig, TrainTeacherTimings, @@ -198,7 +202,12 @@ class WorkerConfig(BaseModel): rollout_steps_per_sft: int = 1 sft_loss_cfg: CELossConfig = CELossConfig() - def build(self, placement_group: "PlacementGroup"): + def build( + self, + placement_group: PlacementGroup, + advantage_estimator: AdvantageEstimator | None = None, + distillation: DistillationTrainerAdapter | None = None, + ): """Build training workers and controller from this config and placement group.""" # import here to avoid circular import @@ -216,7 +225,11 @@ def build(self, placement_group: "PlacementGroup"): )(TrainingWorker) train_workers, _ = AutoAcceleratorWorkers.from_placement_group(TrainingWorkerCls, self, placement_group) ray.wait([w.ready.remote() for w in train_workers]) - return TrainingController(workers=train_workers) + return TrainingController( + workers=train_workers, + advantage_estimator=advantage_estimator, + distillation=distillation, + ) class WorkerInputItem(TypedDict): @@ -244,7 +257,7 @@ class WorkerLogItem(TypedDict): rollout_entropy: NotRequired[float] mismatch_metrics: NotRequired[dict[str, float]] rollout_is_metrics: NotRequired[dict[str, float]] - train_metrics: List[WorkerTrainLogItem] + train_metrics: list[WorkerTrainLogItem] sft_train_metrics: NotRequired[dict[str, float]] diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 1bc2dbc6d0..5532dc1beb 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -1,7 +1,6 @@ import asyncio import json import os -import random import re import time from dataclasses import asdict, dataclass @@ -9,7 +8,6 @@ from shutil import rmtree from typing import Any, List, cast -import numpy as np import ray import torch from mmengine.dist import get_rank @@ -19,8 +17,7 @@ from transformers import AutoTokenizer, PreTrainedTokenizer, PreTrainedTokenizerFast from xtuner.v1._writer import get_writer -from xtuner.v1.data_proto.rl_data import RolloutState, Status -from xtuner.v1.data_proto.sequence_context import SequenceContext +from xtuner.v1.data_proto.rl_data import RolloutState, is_valid_for_training from xtuner.v1.patch import patch_default_save_plan from xtuner.v1.rl.advantage import BaseAdvantageConfig, GRPOAdvantageConfig from xtuner.v1.rl.agent_loop_manager import ( @@ -69,13 +66,6 @@ DEVICE_MODULE = get_torch_device_module() -def _to_cpu_tensor(value: np.ndarray | None, *, dtype: torch.dtype | None = None) -> torch.Tensor | None: - if value is None: - return None - assert isinstance(value, np.ndarray), f"Expected np.ndarray, got {type(value)}" - return torch.as_tensor(value, dtype=dtype, device="cpu") - - def _agent_loop_manager_requires_rollout_proxy( cfg: AgentLoopManagerConfig | DisaggAgentLoopManagerConfig | None, ) -> bool: @@ -275,101 +265,6 @@ def to_scalars(self) -> dict[str, float]: return {f"throughput/{key}": value for key, value in asdict(self).items()} -def get_train_seq_ctx( - input_ids: torch.LongTensor, - position_ids: np.ndarray | None = None, - multimodal_train_info: dict | None = None, - len_response_ids: int = 0, -): - seq_ctx = SequenceContext.from_input_ids((input_ids,), device="cpu") - position_ids = _to_cpu_tensor(position_ids, dtype=torch.long) - if position_ids is not None and len(position_ids.shape) == 3: - # Match get_rope_index_3: response text continues from a single global - # max(T, H, W), not per-axis maxima. Per-axis max diverges when the - # prompt ends on image tokens (T≈0 while H/W are large). - max_value = position_ids.amax() - response_position_ids = ( - ( - torch.arange( - 1, - len_response_ids + 1, - device=position_ids.device, - dtype=position_ids.dtype, - ) - + max_value - ) - .view(1, 1, -1) - .expand(3, 1, -1) - ) - position_ids = torch.cat([position_ids, response_position_ids], dim=-1) - seq_ctx.position_ids = position_ids # type: ignore[assignment] - assert position_ids.size(-1) == input_ids.size(-1) - - if multimodal_train_info: - seq_ctx.pixel_values = multimodal_train_info.get("pixel_values") - seq_ctx.image_grid_thw = _to_cpu_tensor(multimodal_train_info.get("image_grid_thw"), dtype=torch.long) - num_img_tokens = multimodal_train_info.get("num_img_tokens") - if num_img_tokens is not None: - seq_ctx.num_img_tokens = [num_img_tokens] - return seq_ctx - - -def is_valid_for_training(group_data_items: list[RolloutState], logger) -> bool: - """Checks if a group of rollout states is valid for a training step. - - Args: - group_data_items: A list of RolloutState objects. - - Returns: - True if the group is valid, False otherwise. - - NOTE: Why this check is needed: - - For system fault tolerance, this check is performed at rollout / dataflow - time, but we still do it here to ensure training data integrity. - - 'filtered'/'failed': These items are fundamentally broken or incomplete and - should not be used for training. - - 'aborted': These items represent rollouts that were stopped - prematurely. Using such partial data could lead the model to learn - undesirable behaviors (e.g., stopping generation too early). - - Empty response/response_ids: The model's generated response is the core - of the training data for RL algorithms like PPO. If the response is - missing, there is nothing to compute rewards on or to train the model with. - """ - is_abort = any(item.status == Status.ABORTED for item in group_data_items) - is_filtered = any(item.status == Status.FILTERED for item in group_data_items) - is_failed = any(item.status == Status.FAILED for item in group_data_items) - if is_filtered or is_failed or is_abort: - logger.warning( - f"Invalid dataflow group found during training, rollout state filtered: {is_filtered}, failed: {is_failed}, aborted: {is_abort}." - ) - return False - for item in group_data_items: - if item.input_ids is not None: - input_ids_valid = len(item.input_ids) > 1 - labels_valid = item.labels is not None and len(item.labels) == len(item.input_ids) - logprobs_valid = item.logprobs is None or len(item.logprobs) == len(item.input_ids) - if not input_ids_valid or not labels_valid or not logprobs_valid: - logger.warning( - "Invalid dataflow item found during training: input_ids, labels, and logprobs lengths mismatch." - ) - return False - continue - - response_valid = item.response is not None and len(item.response) > 0 - ids_valid = item.response_ids is not None and len(item.response_ids) > 0 - if not ids_valid: - # NOTE: `response_ids` is the critical field for token-in-token-out mode, so we ensure it's not empty. - logger.warning( - "Invalid dataflow item found during training: no response or response_ids and skip this item." - ) - return False - if not response_valid: - # NOTE: check valid response string for judger inputs - logger.warning("Invalid dataflow item found during training: empty response string and skip this item.") - return False - return True - - def _validate_sync_intervals( sync_weights_interval: int, checkpoint_interval: int | None, @@ -1059,20 +954,13 @@ def _train_one_batch( self.train_controller.onload(target="all") self.logger.info("Training controller loaded") - with timer("prepare_data", step_timer_dict): - data_batches, data_info = self._prepare_train_data( - train_batch, - self._train_worker_cfg.pack_max_length, - raw_rewards_sum=raw_rewards_sum, - raw_rewards_count=raw_rewards_count, - ) - self.logger.info(f"Prepared {len(data_batches)} training data batches") - with timer("training", step_timer_dict): - workers_log_item: list[WorkerLogItem] = self.train_controller.fit( - data_batches, + workers_log_item, data_info = self.train_controller.fit( + train_batch, pack_max_length=self._train_worker_cfg.pack_max_length, rollout_idx=train_step, + raw_rewards_sum=raw_rewards_sum, + raw_rewards_count=raw_rewards_count, ) self._release_trace_sessions_after_train_batch(train_batch) @@ -1137,291 +1025,6 @@ def _load_debug_rollout_batch(self, train_step: int) -> list[list[RolloutState]] self.logger.info(f"Loaded debug rollout batch for step {train_step} from {debug_file}") return cast(list[list[RolloutState]], train_batch) - # TODO: simplify with Packer.pack_pad_dispatch() - def _prepare_train_data( - self, - data_groups: list[list[RolloutState]], - pack_max_length: int, - raw_rewards_sum: float = 0.0, - raw_rewards_count: int = 0, - ): - has_agentic_data = any(state.input_ids is not None for group in data_groups for state in group) - has_3d_reasoning_data = any( - state.input_ids is None and state.position_ids is not None and state.position_ids.ndim == 3 - for group in data_groups - for state in group - ) - use_3d_agentic_position_ids = has_agentic_data and has_3d_reasoning_data - - # Per-session rewards for distribution metrics. Agentic sessions may split into several - # trainable segments that share one reward; counting that reward once per session keeps - # rewards/* from being weighted by segment count. - cluster_rewards_list: list[float] = [] - distillation_reward_observations: list[tuple[RolloutState, float]] = [] - advantages_list: list[float] = [] - prompt_len_list = [] - response_len_list = [] - tool_turns_list: list[int] = [] - training_tokens = 0 - training_samples = 0 - - data_batches = [] - - for j, group in enumerate(data_groups): - if not is_valid_for_training(group, self.logger): - self.logger.error(f"Skip one data group {group} due to rollout failed or empty response.") - continue - training_samples += len(group) - - task_adv_weight = self._distillation.task_adv_weight - prompt_ids = None - if any(data.input_ids is None for data in group): - is_vlm_model = "train_prompt_ids" in group[0].extra_fields - if is_vlm_model: - # TODO(hha): VLM, 不好的设计,后续要去掉 - prompt_ids = group[0].extra_fields["train_prompt_ids"] - else: - prompt_ids = group[0].prompt_ids - assert prompt_ids is not None and len(prompt_ids) > 0, ( - f"Prompt ids cannot be None or empty in data: {group[0]}" - ) - # Collect rewards independently from task-advantage computation. Pure OPD may omit - # rewards entirely; when rewards are present they remain useful observability signals. - cluster_index_by_key: dict[Any, int] = {} - cluster_rewards: list[float] = [] - cluster_representatives: list[RolloutState] = [] - sample_cluster_indices: list[int] = [] - for data in group: - # 有可能有重复,但是没有其他更好办法 - turns = data.extra_fields.get("agent_tool_turns") - if isinstance(turns, int): - tool_turns_list.append(turns) - if data.reward is None or "score" not in data.reward: - if task_adv_weight > 0: - raise ValueError( - f"Reward is missing or does not contain 'score' key in data: {data}, " - f"but task_adv_weight={task_adv_weight} > 0" - ) - else: - continue - reward = float(data.reward["score"]) - # session_id is only set by agentic loops / XTUNER_DETERMINISTIC; plain RL falls back - # to rollout_id, which the sampler always assigns. Segments of one session share a key. - cluster_key = data.session_id if data.session_id is not None else data.rollout_id - cluster_index = cluster_index_by_key.get(cluster_key) - if cluster_index is None: - cluster_index = len(cluster_rewards) - cluster_index_by_key[cluster_key] = cluster_index - cluster_rewards.append(reward) - cluster_representatives.append(data) - sample_cluster_indices.append(cluster_index) - - cluster_rewards_list.extend(cluster_rewards) - distillation_reward_observations.extend(zip(cluster_representatives, cluster_rewards)) - - if task_adv_weight == 0: - sample_advantages = [0.0] * len(group) - else: - rewards_tensor = torch.tensor(cluster_rewards, dtype=torch.float32) - cluster_advantages = self._advantage_estimator.compute(rewards_tensor, cluster_representatives) - sample_advantages = [ - cluster_advantages[cluster_index].item() for cluster_index in sample_cluster_indices - ] - - prompt_repeat_k = len(group) - for i in range(prompt_repeat_k): - if group[i].input_ids is not None: - raw_input_ids = cast(list[int], group[i].input_ids) - labels = cast(list[int] | None, group[i].labels) - assert labels is not None, f"Labels cannot be None when input_ids is provided: {group[i]}" - assert len(raw_input_ids) == len(labels), ( - f"{len(raw_input_ids)} vs {len(labels)}, data: {group[i]}" - ) - - input_logprobs = group[i].logprobs - if input_logprobs is not None: - assert len(input_logprobs) == len(raw_input_ids), ( - f"{len(input_logprobs)} vs {len(raw_input_ids)}, data: {group[i]}" - ) - rollout_logprobs: torch.Tensor | None = torch.tensor( - input_logprobs[1:], dtype=torch.float32 - ).unsqueeze(0) - else: - raise ValueError(f"Logprobs cannot be None when input_ids is provided: {group[i]}") - - input_ids = raw_input_ids[:-1] - shifted_labels = labels[1:] - teacher_fields = self._distillation.rollout_teacher_targets( - group[i], - shifted_labels=shifted_labels, - ) - prompt_len = sum(label == -100 for label in shifted_labels) - response_len = len(shifted_labels) - prompt_len - prompt_len_list.append(prompt_len) - response_len_list.append(response_len) - - advatnages_val = sample_advantages[i] - actual_advantages = [0.0 if label == -100 else advatnages_val for label in shifted_labels] - advantages_list.extend(advatnages_val for label in shifted_labels if label != -100) - - assert len(input_ids) <= pack_max_length, f"{len(input_ids)} vs {pack_max_length}" - training_tokens += len(input_ids) - input_ids_t = cast(torch.LongTensor, torch.tensor(input_ids, dtype=torch.int64).unsqueeze(0)) - shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) - - if rollout_logprobs is not None: - assert rollout_logprobs.size() == shifted_labels_t.size(), ( - f"{rollout_logprobs.size()} vs {shifted_labels_t.size()}" - ) - - position_ids = group[i].position_ids - if use_3d_agentic_position_ids: - seq_len = input_ids_t.size(-1) - text_position_ids = np.arange(seq_len, dtype=np.int64).reshape(1, 1, -1) - # Mixed agentic/VLM batches use three-axis MRoPE position IDs. Repeat the - # normal text positions on every axis so agentic samples have the same rank - # as the VLM samples when their sequence contexts are packed together. - position_ids = np.broadcast_to(text_position_ids, (3, 1, seq_len)).copy() - multimodal_train_info = group[i].mm_info - multi_info_cast = cast(dict | None, multimodal_train_info) - seq_ctx = get_train_seq_ctx(input_ids_t, position_ids, multi_info_cast) - - data_dict = { - "seq_ctx": seq_ctx, - "shifted_labels": shifted_labels_t, - "advantage": actual_advantages, - "rollout_logprobs": rollout_logprobs, - } - data_dict.update(teacher_fields) - - seq_ctx.rollout_routed_experts = group[i].routed_experts - data_batches.append(data_dict) - continue - - item = group[i].response - response_logprobs: list[float] | None = None - - response_ids: List[int] = [] - assert prompt_ids is not None - if group[i].response_ids is not None: - resp_ids_raw = group[i].response_ids - if isinstance(resp_ids_raw, torch.Tensor): - response_ids = resp_ids_raw.flatten().tolist() - else: - response_ids = cast(List[int], resp_ids_raw) - - response_logprobs = group[i].logprobs - if response_logprobs is not None: - assert len(response_logprobs) == len(response_ids), ( - f"{len(response_logprobs)} vs {len(response_ids)}, data: {group[i]}" - ) - # 只有 response 部分有 logprobs, 需要前面追加 - response_logprobs = [0.0] * (len(prompt_ids) - 1) + response_logprobs - else: - assert item is not None, "response item cannot be None" - response_ids = self.tokenizer(item, return_tensors="pt")["input_ids"].flatten().tolist() - - # 返回的 routed_experts 不包括 eos 的值,实际上也不需要,需要减一 - input_ids = prompt_ids + response_ids[:-1] - - prompt_len_list.append(len(prompt_ids)) - response_len_list.append(len(response_ids)) - - # 根据 response_mask 计算 response_ids 对应的shifted_labels - if group[i].response_mask is None: - response_mask = [1] * len(response_ids) - response_labels = response_ids - else: - assert len(group[i].response_mask) == len(response_ids), ( # type: ignore[arg-type] - f"{len(group[i].response_mask)} vs {len(response_ids)}" # type: ignore[arg-type] - ) - response_mask = cast(list[int], group[i].response_mask) - response_labels = [ - response_id if mask_id != 0 else -100 - for response_id, mask_id in zip(response_ids, response_mask) - ] - shifted_labels = [-100] * (len(prompt_ids) - 1) + response_labels - shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) - - # Keep the advantage layout aligned with input_ids (response excludes EOS). - # Prompt positions predict prompt tokens whose shifted labels are -100, so their - # advantage is 0; the last entry of response_ids (EOS) is only a label, never an input. - advatnages_val = sample_advantages[i] - actual_advantages = [0.0] * (len(prompt_ids) - 1) + [ - 0.0 if mask == 0 else advatnages_val for mask in response_mask - ] - advantages_list.extend(advatnages_val for mask in response_mask if mask != 0) - - assert len(input_ids) <= pack_max_length, f"{len(input_ids)} vs {pack_max_length}" - training_tokens += len(input_ids) - input_ids_t = cast(torch.LongTensor, torch.tensor(input_ids, dtype=torch.int64).unsqueeze(0)) - - if response_logprobs is not None: - rollout_logprobs = torch.tensor(response_logprobs, dtype=torch.float32).unsqueeze(0) - assert rollout_logprobs.size() == shifted_labels_t.size(), ( - f"{rollout_logprobs.size()} vs {shifted_labels_t.size()}" - ) - else: - rollout_logprobs = None - - position_ids = group[i].position_ids - multimodal_train_info = group[i].mm_info - multi_info_cast = cast(dict | None, multimodal_train_info) - seq_ctx = get_train_seq_ctx(input_ids_t, position_ids, multi_info_cast, len(response_ids) - 1) # type: ignore[arg-type] - - data_dict = { - "seq_ctx": seq_ctx, - "shifted_labels": shifted_labels_t, - "advantage": actual_advantages, - "rollout_logprobs": rollout_logprobs, - } - teacher_fields = self._distillation.rollout_teacher_targets( - group[i], - shifted_labels=shifted_labels, - ) - data_dict.update(teacher_fields) - - seq_ctx.rollout_routed_experts = group[i].routed_experts # n,layer*expert - - data_batches.append(data_dict) - if not XTUNER_DETERMINISTIC: - random.shuffle(data_batches) - - # rewards/* report the per-session reward distribution; batch_size/training_samples - # count the valid rollout segments included in training. - rewards_t = torch.tensor(cluster_rewards_list).float() if cluster_rewards_list else torch.tensor([0.0]).float() - advantages_t = torch.tensor(advantages_list).float() if advantages_list else torch.tensor([0.0]).float() - prompt_len_t = torch.tensor(prompt_len_list).float() if prompt_len_list else torch.tensor([0.0]).float() - response_len_t = torch.tensor(response_len_list).float() if response_len_list else torch.tensor([0.0]).float() - - raw_rewards_mean = raw_rewards_sum / raw_rewards_count if raw_rewards_count > 0 else rewards_t.mean().item() - info_dict = { - "batch_size": training_samples, - "training_samples": training_samples, - "training_tokens": training_tokens, - "rewards/mean": rewards_t.mean().item(), - "rewards/min": rewards_t.min().item(), - "rewards/max": rewards_t.max().item(), - "raw_rewards/mean": raw_rewards_mean, - "advantages/mean": advantages_t.mean().item(), - "advantages/min": advantages_t.min().item(), - "advantages/max": advantages_t.max().item(), - "response_len/mean": response_len_t.mean().item(), - "response_len/min": response_len_t.min().item(), - "response_len/max": response_len_t.max().item(), - "response_len/std": response_len_t.std().item(), - "prompt_len/mean": prompt_len_t.mean().item(), - "prompt_len/min": prompt_len_t.min().item(), - "prompt_len/max": prompt_len_t.max().item(), - } - info_dict.update(self._distillation.reward_scalars(distillation_reward_observations)) - if tool_turns_list: - tool_turns_t = torch.tensor(tool_turns_list, dtype=torch.float32) - info_dict["tool_turns/mean"] = tool_turns_t.mean().item() - info_dict["tool_turns/min"] = float(tool_turns_t.min().item()) - info_dict["tool_turns/max"] = float(tool_turns_t.max().item()) - return data_batches, info_dict - def _compute_benchmark_metrics( self, data_info: dict[str, float], @@ -1777,7 +1380,11 @@ def __init__(self, cfg: RLColocateTrainerConfig): return - self.train_controller = self._train_worker_cfg.build(self._pg) + self.train_controller = self._train_worker_cfg.build( + self._pg, + advantage_estimator=self._advantage_estimator, + distillation=self._distillation, + ) checkpoint_path = self._load_checkpoint_cfg.checkpoint_path if checkpoint_path is not None: @@ -2041,7 +1648,11 @@ def __init__(self, cfg: RLDisaggregatedTrainerConfig): self._cpu_resource_manager = CPUResourceManager([self._train_pg, self._rollout_pg]) self._cpu_resource_manager.log_initial_snapshot() set_cpu_resource_manager(self._cpu_resource_manager) - self.train_controller = self._train_worker_cfg.build(self._train_pg) + self.train_controller = self._train_worker_cfg.build( + self._train_pg, + advantage_estimator=self._advantage_estimator, + distillation=self._distillation, + ) self.rollout_controller = self._rollout_config.build(self._rollout_pg) if self._rollout_config.weight_transport_type != "nccl": From 066bd24e379acc7e4efc5a4ed3315dd138880801 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Mon, 21 Sep 2026 06:18:44 +0000 Subject: [PATCH 2/4] [Refactor] Unify response_ids suffix convention and phase-split rollout conversion - response_ids now denotes the contiguous suffix of input_ids after the prompt (env/tool tokens included) in every loop; localhost/sandbox export input_ids[len(prompt_ids):] so response_model_steps stay aligned with the full response region. - Token staleness baking becomes branch-free: effective mask = semantic mask (labels != -100 on the suffix) * per-token staleness mask; the zero-prompt suffix state becomes eligible and agent_loop_type plus AGENTIC_AGENT_LOOP_TYPES are removed. - data_info stats: prompt_len reports the original prompt length and response_len the supervised (LLM-generated) token count; env/tool injected tokens count in neither. - _rollout_groups_to_colate_items is phase-split into session reward clustering, group advantage estimation, per-state ColateItem conversion and data_info summarization. --- docs/zh_cn/rl/advanced_tutorial/agent_loop.md | 8 +- .../rl/test_multi_task_agent_loop_manager.py | 42 +- tests/rl/test_prepare_train_data.py | 24 +- .../test_rl_colocate_trainer_integration.py | 2 - tests/rl/test_staleness_policy.py | 102 ++++- xtuner/v1/data_proto/rl_data.py | 69 ++-- xtuner/v1/rl/agent_loop/agent_loop.py | 2 - .../agent_in_localhost_loop.py | 9 +- .../agent_in_sandbox_loop.py | 11 +- xtuner/v1/rl/trainer/controller.py | 361 +++++++++++------- 10 files changed, 404 insertions(+), 226 deletions(-) diff --git a/docs/zh_cn/rl/advanced_tutorial/agent_loop.md b/docs/zh_cn/rl/advanced_tutorial/agent_loop.md index e87980d240..5335188c96 100644 --- a/docs/zh_cn/rl/advanced_tutorial/agent_loop.md +++ b/docs/zh_cn/rl/advanced_tutorial/agent_loop.md @@ -92,9 +92,9 @@ AgentLoop 返回的 `RolloutState` 如果要进入训练,至少需要满足: - `logprobs`:长度必须等于 `len(response_ids)`。 -labels 是唯一的监督载体。`AgentLoop.generate_group()` 会把产出 loop 的类名写入 `agent_loop_type`,并在末尾调用 `canonicalize_train_fields()`,为 prompt+response 型样本构造全序列 `input_ids`/`labels`/`logprobs`(三者等长、未 shift)。这也是自定义 AgentLoop 最容易出错的地方:工具返回、环境反馈、系统插入内容等不是模型生成的 token,不参与训练——对应 label 直接写 `-100`,`logprobs` 填 `0.0`。prompt+response 型样本不需要手动构造这些字段(基类会兜底);自行组装全序列的多轮 loop 则必须自己把语义洞烙进 labels。训练侧只做 shift 和 advantage 计算,不处理任何掩码。 +labels 是唯一的监督载体。`AgentLoop.generate_group()` 会在末尾调用 `canonicalize_train_fields()`,为 prompt+response 型样本构造全序列 `input_ids`/`labels`/`logprobs`(三者等长、未 shift)。这也是自定义 AgentLoop 最容易出错的地方:工具返回、环境反馈、系统插入内容等不是模型生成的 token,不参与训练——对应 label 直接写 `-100`,`logprobs` 填 `0.0`。prompt+response 型样本不需要手动构造这些字段(基类会兜底);自行组装全序列的多轮 loop 则必须自己把语义洞烙进 labels。训练侧只做 shift 和 advantage 计算,不处理任何掩码。 -agentic 全序列 loop(如 `AgentInLocalhostLoop`、`AgentInSandboxLoop`)的类名需要登记到 `xtuner.v1.data_proto.rl_data.AGENTIC_AGENT_LOOP_TYPES`,其样本才会被 token 级 staleness 排除;未登记的自定义 loop 默认按 prompt+response(reasoning)样本处理,`agent_loop_type` 为 `None`(未记录,如手工构造的样本)时同样按 prompt+response 处理。 +`response_ids` 统一约定为 `input_ids` 去掉 prompt 前缀的连续后缀(`len(response_ids) == len(input_ids) - len(prompt_ids)`,环境插入或工具返回的 token 也包含在内);`response_model_steps` 与之等长,记录每位 token 的来源模型版本。token 级 staleness 据此对 prompt+response 与 agentic 全序列样本统一生效:逐 token 判定新鲜度,过期 token 的 label 烙成 `-100`,语义洞位(label 为 `-100`)不受影响。 ## SingleTurnAgentLoop @@ -148,7 +148,7 @@ agent_loop_config = SingleTurnAgentLoopConfig( 自定义 AgentLoop 通常需要做四件事: 1. 继承 `AgentLoop`,实现 `generate_sample()`。 -2. 在 `generate_sample()` 中维护 `tokens`、`sample_params`、`response_ids`、`response`、`logprobs`、`status`,不要在这里调用外部 Judger。prompt+response 型样本的训练字段(`input_ids`/`labels`)由基类 `canonicalize_train_fields()` 兜底构造;自行组装全序列的 loop 需自己写 `input_ids`/`labels`,并把语义洞直接烙进 labels。 +2. 在 `generate_sample()` 中维护 `tokens`、`sample_params`、`response_ids`、`response`、`logprobs`、`status`,不要在这里调用外部 Judger。prompt+response 型样本的训练字段(`input_ids`/`labels`)由基类 `canonicalize_train_fields()` 兜底构造;自行组装全序列的 loop 需自己写 `input_ids`/`labels`,并把语义洞直接烙进 labels,同时把 `response_ids` 写成 `input_ids` 去掉 prompt 前缀的连续后缀(含环境 token)。 3. 若覆盖 `generate_group()`,在其中显式编排 Teacher、Judger 和组级过滤;没有 validity check 时应让完成的样本尽早进入 Teacher,有 validity check 时应只把过滤通过的完整组发送给 Teacher。 4. 继承 `AgentLoopConfig`,实现 `build_local()`,这样才能接入 `TaskSpecConfig.agent_loop_config`,并复用 Ray actor 构建逻辑。 @@ -350,7 +350,7 @@ agent_loop_manager_cfg = AgentLoopManagerConfig( - 每次调用 `rollout_ctl.generate.remote()` 前是否设置了本轮 `sample_params`。 - 返回训练前,`response_ids`、`response`、`logprobs` 是否完整且长度一致;自组装全序列的 loop 还需保证 `input_ids`/`labels` 齐全且等长。 - 非模型生成 token 的 label 是否已写 `-100`(prompt+response 型样本由基类 canonicalize 兜底;自组装 loop 自行烙入)。 -- 若是 agentic 全序列 loop,类名是否已登记到 `AGENTIC_AGENT_LOOP_TYPES`。 +- `response_ids` 是否为 `input_ids` 去掉 prompt 前缀的连续后缀(`len(response_ids) == len(input_ids) - len(prompt_ids)`)。 - 需要 Judger 时,是否通过 `self.run_judger(...)` 调用打分,以复用 pause/cancel 处理。 - 是否保证最终有 `reward["score"]`。 - 若使用 async partial rollout,是否正确处理 `enable_partial_rollout` 和历史 response 合并。 diff --git a/tests/rl/test_multi_task_agent_loop_manager.py b/tests/rl/test_multi_task_agent_loop_manager.py index 35cdf7f500..549ded8d90 100644 --- a/tests/rl/test_multi_task_agent_loop_manager.py +++ b/tests/rl/test_multi_task_agent_loop_manager.py @@ -349,7 +349,8 @@ async def test_take_train_batch_applies_token_staleness_mask(self): self.assertEqual(result.rollout_states[0][0].labels, [-100, -100, -100, 4]) self.assertEqual(replay_buffer.task_token_stale_threshold_calls, [{"task": 4}]) - async def test_take_train_batch_skips_agentic_token_staleness_mask(self): + async def test_take_train_batch_skips_states_without_response_ids(self): + # response_ids 缺失的样本没有可对齐的 response 段,不参与 token staleness 烙制。 state = RolloutState( rollout_id=1, group_id=1, @@ -357,7 +358,6 @@ async def test_take_train_batch_skips_agentic_token_staleness_mask(self): input_ids=[1, 2], labels=[-100, 2], logprobs=[0.0, -0.1], - agent_loop_type="AgentInLocalhostLoop", status=Status.COMPLETED, ) strategy = _FakeProduceStrategy(token_stale_threshold=4) @@ -383,6 +383,44 @@ async def test_take_train_batch_skips_agentic_token_staleness_mask(self): self.assertEqual(result.rollout_states[0][0].labels, [-100, 2]) + async def test_take_train_batch_bakes_full_sequence_agentic_per_token_staleness(self): + # agentic 全序列形态(response_ids 为 input_ids 去掉 prompt 前缀的后缀,含环境 token) + # 同样按 token 级 staleness 烙制:逐 token 按 response_model_steps 判定,仅过期 token + # 的监督位被清零。 + state = RolloutState( + rollout_id=1, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + prompt_ids=[1, 2], + response_ids=[30, -1, 40], + response_model_steps=[0, 4, 4], + input_ids=[1, 2, 30, -1, 40], + labels=[-100, -100, 30, -100, 40], + status=Status.COMPLETED, + ) + strategy = _FakeProduceStrategy(token_stale_threshold=4) + manager = AgentLoopManager( + task_runners=[ + _TaskRunner( + task_name="task", + agent_loop=_fake_agent_loop(), + produce_strategy=strategy, + sampler=_FakeSampler(), + weight=1.0, + order=0, + ) + ], + replay_buffer=_FakeReplayBuffer( + rollout_states_by_task={"task": [[state]]}, + leftover_counts={}, + ), + rollout_controller=_fake_rollout_controller(), + ) + + result = await manager.produce_batch(batch_size=1, train_step=5, model_step=4) + + self.assertEqual(result.rollout_states[0][0].labels, [-100, -100, -100, -100, 40]) + async def test_take_train_batch_sync_path_leaves_labels_untouched(self): # 同步路径(无 token staleness)零操作:labels 在生成期已是最终态(语义洞已烙入)。 state = RolloutState( diff --git a/tests/rl/test_prepare_train_data.py b/tests/rl/test_prepare_train_data.py index f9cf161a87..e018ba2806 100644 --- a/tests/rl/test_prepare_train_data.py +++ b/tests/rl/test_prepare_train_data.py @@ -4,10 +4,9 @@ rollout backend: - loop 侧 ``canonicalize_train_fields``:基类为 prompt+response 型样本构造统一全序列字段 - (``input_ids``/``labels``/``logprobs`` 等长、未 shift;``agent_loop_type`` 由 generate_group - 写入产出 loop 的类名);localhost/sandbox 各自覆写为 - trace 全序列字段的校验(失败置 FAILED,不上抛)。 -- controller 侧 ``TrainingController._convert_rollout_groups``:消费已 canonicalize 且 labels 已定稿 + (``input_ids``/``labels``/``logprobs`` 等长、未 shift); + localhost/sandbox 各自覆写为 trace 全序列字段的校验(失败置 FAILED,不上抛)。 +- controller 侧 ``TrainingController._rollout_groups_to_colate_items``:消费已 canonicalize 且 labels 已定稿 (语义洞直接烙在 labels 中)的状态,完成 shift、张量化、组级 advantage、seq_ctx 构造与 ``data_info`` 统计。 @@ -319,7 +318,7 @@ def test_non_completed_state_is_skipped(self): class TestConvertRolloutGroups(unittest.TestCase): - """TrainingController._convert_rollout_groups 合同:shift、advantage、张量与统计。""" + """TrainingController._rollout_groups_to_colate_items 合同:shift、advantage、张量与统计。""" def _build_controller(self, advantages: list[float], task_adv_weight: float = 1.0) -> TrainingController: controller = TrainingController.__new__(TrainingController) @@ -331,7 +330,7 @@ def _build_controller(self, advantages: list[float], task_adv_weight: float = 1. def _convert(self, controller, data_groups, pack_max_length=128): with patch("xtuner.v1.rl.trainer.controller.XTUNER_DETERMINISTIC", True): - return controller._convert_rollout_groups(data_groups, pack_max_length) + return controller._rollout_groups_to_colate_items(data_groups, pack_max_length) def _state( self, @@ -353,7 +352,6 @@ def _state( labels: list[int] | None = None, teacher_tokens: list[int] | list[list[int]] | None = None, teacher_logprobs: list[float] | list[list[float]] | None = None, - agent_loop_type: str | None = None, ) -> RolloutState: resolved_prompt_ids = prompt_ids if prompt_ids is not None else [10, 11, 12] resolved_response_ids = response_ids if response_ids is not None else [20, 21, 22] @@ -365,7 +363,6 @@ def _state( response=response, response_ids=resolved_response_ids, logprobs=logprobs, - agent_loop_type=agent_loop_type, reward=reward if reward is not None else {"score": 1.0}, status=status, finish_reason="stop" if status == Status.COMPLETED else "error", @@ -441,7 +438,8 @@ def test_text_path_builds_shifted_training_tensors(self): self.assertEqual(info["training_samples"], 1) self.assertEqual(info["training_tokens"], 5) self.assertEqual(info["rewards/mean"], 1.0) - self.assertEqual(info["response_len/mean"], 3.0) + # response_len 从 labels 监督位推导(labels 是唯一监督载体):语义洞不计入。 + self.assertEqual(info["response_len/mean"], 2.0) self.assertEqual(info["prompt_len/mean"], 3.0) def test_controller_consumes_labels_as_final_supervision(self): @@ -490,7 +488,6 @@ def test_advantage_stats_count_only_loss_active_tokens(self): labels=[-100, -100, 40, -100, 42], logprobs=[0.0, -0.1, -0.2, -0.3, -0.4], reward={"score": -1.0}, - agent_loop_type="AgentInLocalhostLoop", ) _, info = self._convert(controller, [[plain, agentic]]) @@ -550,7 +547,6 @@ def test_get_train_seq_ctx_mrope_continues_from_global_amax(self): seq_ctx = get_train_seq_ctx( cast(torch.LongTensor, full_ids), prompt_position_ids.numpy(), - len_response_ids=len(response_ids), ) rl_position_ids = seq_ctx.position_ids assert rl_position_ids is not None @@ -572,7 +568,6 @@ def test_mixed_agentic_and_vlm_reasoning_use_3d_position_ids(self): input_ids=[30, 31, 40, 41, 42], labels=[-100, -100, 40, 41, 42], logprobs=[0.0, -0.1, -0.2, -0.3, -0.4], - agent_loop_type="AgentInLocalhostLoop", ) data_batches, _ = self._convert(controller, [[reasoning_state], [agentic_state]]) @@ -605,7 +600,6 @@ def test_agentic_topk_targets_include_token_ids_and_logprobs(self): teacher_tokens=[[100, 101], [102, 103], [104, 105]], teacher_logprobs=[[-0.5, -0.6], [-0.7, -0.8], [-0.9, -1.0]], extra_fields={"origin_data_source": "agent_math"}, - agent_loop_type="AgentInLocalhostLoop", ) data_batches, _ = self._convert(controller, [[state]]) @@ -705,7 +699,6 @@ def test_sampled_token_targets_align_for_plain_and_agentic_rollouts(self): teacher_tokens=[40, 41, 42], teacher_logprobs=[-1.1, -1.2, -1.3], extra_fields={"origin_data_source": "agent_math"}, - agent_loop_type="AgentInLocalhostLoop", ) data_batches, _ = self._convert(controller, [[plain_state], [agentic_state]]) @@ -763,7 +756,6 @@ def test_group_with_mismatched_full_sequence_logprobs_is_skipped(self): input_ids=[30, 31, 40, 41, 42], labels=[-100, -100, 40, 41, 42], logprobs=[0.0, -0.1, -0.2, -0.3], - agent_loop_type="AgentInLocalhostLoop", ) data_batches, info = self._convert(controller, [[state]]) @@ -889,7 +881,7 @@ def _build_data_groups(self) -> list[list[RolloutState]]: def _convert(self, controller, data_groups, pack_max_length: int): with patch("xtuner.v1.rl.trainer.controller.XTUNER_DETERMINISTIC", True): - return controller._convert_rollout_groups(data_groups, pack_max_length) + return controller._rollout_groups_to_colate_items(data_groups, pack_max_length) def _assert_real_token_alignment( self, samples: list[RolloutState], packed: dict, total_len: int, packed_len: int diff --git a/tests/rl/test_rl_colocate_trainer_integration.py b/tests/rl/test_rl_colocate_trainer_integration.py index 8de969647f..b4c9580f6f 100644 --- a/tests/rl/test_rl_colocate_trainer_integration.py +++ b/tests/rl/test_rl_colocate_trainer_integration.py @@ -255,7 +255,6 @@ def test_rl_train_with_sft(self): # - input_ids keeps the whole prompt+response sequence # - labels supervise the full response; TrainingController applies the # one-position shift at conversion time - # - agent_loop_type records the producing loop's class name group.append(RolloutState( rollout_id=group_idx * len(response_list) + i, group_id=group_idx, @@ -263,7 +262,6 @@ def test_rl_train_with_sft(self): prompt_ids=prompt_ids, response=response, response_ids=response_ids, - agent_loop_type="SingleTurnAgentLoop", reward={"score": group_rewards[i]}, status=Status.COMPLETED, input_ids=prompt_ids + response_ids, diff --git a/tests/rl/test_staleness_policy.py b/tests/rl/test_staleness_policy.py index 7f70db9edf..8a8785b864 100644 --- a/tests/rl/test_staleness_policy.py +++ b/tests/rl/test_staleness_policy.py @@ -72,7 +72,7 @@ def test_async_strategies_precompute_token_stale_threshold(self): class TestTokenStalenessMask(unittest.TestCase): - """Token 级 staleness mask 的阈值与语义监督(labels 尾段)行为。""" + """Token 级 staleness mask 的阈值与语义监督行为(labels 尾段 / agentic 全序列两种形态)。""" def test_token_staleness_threshold_can_be_relaxed(self): # token threshold 放宽一个同步周期后,旧周期 token 应从 masked 变为可训练。 @@ -147,21 +147,94 @@ def test_reset_clears_canonical_train_fields(self): self.assertIsNone(state.input_ids) self.assertIsNone(state.labels) - def test_agentic_loop_type_is_excluded(self): - # agentic loop 类型(localhost/sandbox)产出的全序列样本不参与 token staleness,整组排除。 - agentic = self._state(response_model_steps=None, agentic=True) - agentic.input_ids = [1, 2, 3, 4] - agentic.labels = [-100, -100, 3, 4] - agentic.logprobs = [0.0, 0.0, -0.1, -0.2] + def test_full_sequence_agentic_stale_clears_all_supervised_labels(self): + # agentic 全序列形态:response_ids 是 input_ids 去掉 prompt 前缀的连续后缀(含环境 + # token)。steps 由 replay buffer 回填为单一 model_step 时全部过期,全部监督位清零。 + state = RolloutState( + rollout_id=1, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + prompt_ids=[1, 2], + response_ids=[30, -1, 40, -1, 32], + response_model_steps=[0, 0, 0, 0, 0], + input_ids=[1, 2, 30, -1, 40, -1, 32], + labels=[-100, -100, 30, -100, 40, -100, 32], + ) masks = calculate_group_effective_response_masks( - [agentic], + [state], + current_train_step=5, + token_stale_threshold=4, + ) + + self.assertEqual(masks, [[0, 0, 0, 0, 0]]) + self.assertEqual(state.labels, [-100] * 7) + + def test_full_sequence_agentic_per_token_stale_clears_only_stale_tokens(self): + # agentic 全序列形态同样按 token 级 staleness 判定:仅过期 token 的监督位被清零, + # 新鲜 token 原样保留,环境 token 的洞位(semantic=0)不受 staleness 影响。 + state = RolloutState( + rollout_id=1, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + prompt_ids=[1, 2], + response_ids=[30, -1, 40, -1, 32], + response_model_steps=[0, 4, 4, 4, 4], + input_ids=[1, 2, 30, -1, 40, -1, 32], + labels=[-100, -100, 30, -100, 40, -100, 32], + ) + + masks = calculate_group_effective_response_masks( + [state], + current_train_step=5, + token_stale_threshold=4, + ) + + self.assertEqual(masks, [[0, 0, 1, 0, 1]]) + self.assertEqual(state.labels, [-100, -100, -100, -100, 40, -100, 32]) + + def test_full_sequence_agentic_fresh_keeps_labels(self): + # agentic 全序列形态且 token 新鲜:labels 保持不变,掩码即语义掩码(不会触发过期)。 + state = RolloutState( + rollout_id=1, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + prompt_ids=[1, 2], + response_ids=[30, -1, 40, -1, 32], + response_model_steps=[4, 4, 4, 4, 4], + input_ids=[1, 2, 30, -1, 40, -1, 32], + labels=[-100, -100, 30, -100, 40, -100, 32], + ) + + masks = calculate_group_effective_response_masks( + [state], + current_train_step=5, + token_stale_threshold=4, + ) + + self.assertEqual(masks, [[1, 0, 1, 0, 1]]) + self.assertEqual(state.labels, [-100, -100, 30, -100, 40, -100, 32]) + + def test_zero_prompt_suffix_state_is_eligible(self): + # response_ids 允许覆盖整个序列(prompt 前缀长度为 0):len(labels) == + # len(response_ids) 的后缀形态必须仍参与 staleness 烙制,不被 guard 跳过。 + state = RolloutState( + rollout_id=1, + group_id=1, + message=[{"role": "user", "content": "prompt"}], + response_ids=[3, 4], + response_model_steps=[0, 4], + labels=[3, 4], + ) + + masks = calculate_group_effective_response_masks( + [state], current_train_step=5, token_stale_threshold=4, ) - self.assertEqual(masks, [None]) - self.assertEqual(agentic.labels, [-100, -100, 3, 4]) + self.assertEqual(masks, [[0, 1]]) + self.assertEqual(state.labels, [-100, 4]) def test_canonicalized_prompt_response_state_is_still_eligible(self): # canonicalize 后的 prompt+response 样本(reasoning loop 产出)仍参与 staleness。 @@ -183,7 +256,6 @@ def _state( *, response_model_steps: list[int] | None, labels: list[int] | None = None, - agentic: bool = False, ) -> RolloutState: state = RolloutState( rollout_id=1, @@ -193,12 +265,8 @@ def _state( response_ids=[3, 4], response_model_steps=response_model_steps, ) - if agentic: - state.agent_loop_type = "AgentInLocalhostLoop" - else: - # prompt+response 形态:agent_loop_type 保持 None(未记录时默认按 prompt+response 处理), - # labels 为全序列监督。 - state.labels = labels if labels is not None else [-100, -100, 3, 4] + # prompt+response 形态:response_ids 是序列的连续尾段,labels 为全序列监督。 + state.labels = labels if labels is not None else [-100, -100, 3, 4] return state diff --git a/xtuner/v1/data_proto/rl_data.py b/xtuner/v1/data_proto/rl_data.py index cb31b03ca3..512cf6ce71 100644 --- a/xtuner/v1/data_proto/rl_data.py +++ b/xtuner/v1/data_proto/rl_data.py @@ -141,16 +141,13 @@ class RolloutState(BaseModel): routed_experts: np.ndarray | RayObjectRef | list[RayObjectRef] | None = None finish_reason: str | None = None # response_model_steps:记录 response_ids 中每个 token 来自哪个 model_step,与 response_ids 长度相同。 + # 环境/工具等非模型生成位沿用所在 rollout cycle 的 model_step 占位(勿填 0,避免污染 min() 语义)。 response_model_steps: list[int] | None = None # 记录该样本过期程度,即最早生成 token 的模型版本与当前训练步数的差值,数值越大表示越过期。 seq_staleness: int = 0 input_ids: list[int] | None = None labels: list[int] | None = None - # 产出该样本的 AgentLoop 子类类名(generate_group 写入 type(self).__name__)。 - # 结合 AGENTIC_AGENT_LOOP_TYPES 可区分 agentic 全序列样本与 prompt+response(reasoning)样本; - # None(未记录,如手工构造的样本)按 prompt+response 处理。 - agent_loop_type: str | None = None # --- Judger 输出 --- reward: dict[str, Any] | None = None @@ -413,40 +410,39 @@ def refresh_seq_staleness(group: list[RolloutState], current_train_step: int) -> return group -# 产自这些 AgentLoop 子类(类名记录在 RolloutState.agent_loop_type)的样本是 agentic 全序列样本: -# labels 无 prompt 段偏移语义,不参与 token 级 staleness。自定义 agentic loop 需把类名加入此集合。 -AGENTIC_AGENT_LOOP_TYPES: frozenset[str] = frozenset({"AgentInLocalhostLoop", "AgentInSandboxLoop"}) - - def _calculate_effective_response_mask( rollout_state: RolloutState, *, current_train_step: int, token_stale_threshold: int, ) -> list[int]: - """Calculate the response mask after applying token staleness. + """Bake token staleness into labels and return the effective response mask. - The semantic mask is recovered from the labels tail (response segment): supervised - tokens carry their target id, masked-out tokens are ``-100``. + Every loop writes ``response_ids`` under one convention: the contiguous suffix of ``input_ids`` after + the prompt (env-injected or tool tokens included), so response token ``j`` maps to labels position + ``len(labels) - len(response_ids) + j``. The effective mask follows the reasoning-RL rule + ``effective = semantic_mask * token_staleness_mask``: the semantic mask is recovered from + ``labels != -100`` on the response region (semantic holes stay supervised-out regardless of staleness), + staleness is evaluated per token via ``response_model_steps``, and a zero effective mask bakes ``-100`` + into the label in place. Args: - rollout_state (RolloutState): Rollout sample whose response token provenance is evaluated. + rollout_state (RolloutState): Rollout sample whose labels are updated in place. current_train_step (int): Trainer step that will consume the sample. token_stale_threshold (int): Maximum token staleness, measured in trainer steps, allowed for training. Returns: - list[int]: The semantic response mask intersected with the token-staleness mask. + list[int]: The effective mask aligned with ``response_ids``. """ labels = cast(list[int], rollout_state.labels) response_ids = cast(list[int], rollout_state.response_ids) response_model_steps = cast(list[int], rollout_state.response_model_steps) - # semantic mask: 从 labels 的 response 段恢复(-100 即语义上不参与 loss 的 token)。 - # response 段偏移 = len(labels) - len(response_ids),即 prompt 段长度。 + # response_ids 是 input_ids 去掉 prompt 前缀的连续后缀,offset 即 prompt 长度。 offset = len(labels) - len(response_ids) + # semantic_mask: labels 中非 -100 的监督位;token_staleness_mask: 逐 token 新鲜度 + # (语义洞位的 staleness 值无意义,乘上 semantic_mask 后自然归零)。 semantic_mask = [int(label != -100) for label in labels[offset:]] - - # token_staleness_mask: 根据 token 的新鲜程度来 mask token_staleness_mask = [ int(calculate_seq_staleness(response_model_step, current_train_step) < token_stale_threshold) for response_model_step in response_model_steps @@ -455,6 +451,9 @@ def _calculate_effective_response_mask( semantic_mask_value * token_staleness_mask_value for semantic_mask_value, token_staleness_mask_value in zip(semantic_mask, token_staleness_mask) ] + for i, mask_value in enumerate(effective_mask): + if mask_value == 0: + labels[offset + i] = -100 return effective_mask @@ -467,14 +466,12 @@ def calculate_group_effective_response_masks( """Calculate a group's effective masks and bake token staleness into its labels. - For each eligible state, response labels whose token staleness reaches the threshold - are cleared to ``-100`` in place. Clearing only ever extends: staleness grows - monotonically with the trainer step, so repeated calls (e.g. replay-buffer expiry - checks followed by the train-batch bake) converge to the same labels. Each returned - mask is the semantic response mask intersected with the token-staleness mask. - ``None`` means token staleness is disabled or does not apply to that state. Agentic - full-sequence states (``agent_loop_type`` in ``AGENTIC_AGENT_LOOP_TYPES``) receive - ``None``. + For each eligible state, stale supervised labels are cleared to ``-100`` in place. + Clearing only ever extends: staleness grows monotonically with the trainer step, so + repeated calls (e.g. replay-buffer expiry checks followed by the train-batch bake) + converge to the same labels. Each returned mask is the effective response mask + aligned with ``response_ids``. ``None`` means token staleness is disabled or does + not apply to that state (no labels, or no ``response_ids`` to align with). Args: group (list[RolloutState]): Rollout group updated in place. @@ -487,23 +484,17 @@ def calculate_group_effective_response_masks( """ if token_stale_threshold is None: return [None] * len(group) - if any(item.agent_loop_type in AGENTIC_AGENT_LOOP_TYPES for item in group): - return [None] * len(group) masks: list[list[int] | None] = [] for item in group: - if item.labels is None or not item.response_ids or len(item.labels) <= len(item.response_ids): + if item.labels is None or not item.response_ids or len(item.labels) < len(item.response_ids): masks.append(None) continue - effective_mask = _calculate_effective_response_mask( - item, - current_train_step=current_train_step, - token_stale_threshold=token_stale_threshold, + masks.append( + _calculate_effective_response_mask( + item, + current_train_step=current_train_step, + token_stale_threshold=token_stale_threshold, + ) ) - # prompt 段长度由 labels 与 response 段(有效掩码)长度差推导。 - offset = len(item.labels) - len(effective_mask) - for i, mask_value in enumerate(effective_mask): - if mask_value == 0: - item.labels[offset + i] = -100 - masks.append(effective_mask) return masks diff --git a/xtuner/v1/rl/agent_loop/agent_loop.py b/xtuner/v1/rl/agent_loop/agent_loop.py index a25cc579de..b7e8736a5b 100644 --- a/xtuner/v1/rl/agent_loop/agent_loop.py +++ b/xtuner/v1/rl/agent_loop/agent_loop.py @@ -275,8 +275,6 @@ async def generate_one(state: RolloutState) -> RolloutState: state = await self._teacher_scorer.on_sample_ready(state) return state - for state in rollout_state: - state.agent_loop_type = type(self).__name__ group = list(await asyncio.gather(*(create_task(generate_one(state)) for state in rollout_state))) if self.judger is not None and self.enable_batch_judge: if all(sample.status == Status.COMPLETED for sample in group): diff --git a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py index a4ac2da918..00c28e6747 100644 --- a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py +++ b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py @@ -145,7 +145,6 @@ async def generate_one(state: RolloutState) -> RolloutState: tasks: list[asyncio.Task[RolloutState]] = [] for state in rollout_state: - state.agent_loop_type = type(self).__name__ state.sample_params = self.sample_params task = create_task(generate_one(state)) tasks.append(task) @@ -295,9 +294,11 @@ async def _fill_rollout_state(self, rollout_state: RolloutState, item: AgentRoll rollout_state.input_ids = data["input_ids"] rollout_state.labels = data["labels"] - rollout_state.response_ids = [ - token_id for token_id, label in zip(data["input_ids"][1:], data["labels"][1:]) if label != -100 - ] + # Unified response_ids convention: the contiguous suffix of input_ids after the prompt + # (env-injected tokens included), aligned with response_model_steps for per-token + # staleness; also the token count used by rollout throughput logging. + prompt_len = len(rollout_state.prompt_ids or []) + rollout_state.response_ids = list(data["input_ids"][prompt_len:]) rollout_state.logprobs = data["logprobs"] rollout_state.routed_experts = data["routed_experts"] content = response_message.get("content") diff --git a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py index f40e2ea8af..19ea00513b 100644 --- a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py +++ b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py @@ -240,7 +240,6 @@ async def generate_one(state: RolloutState) -> list[RolloutState]: pending_tasks = [] for state in rollout_state: - state.agent_loop_type = type(self).__name__ state.sample_params = self.sample_params task = create_task(generate_one(state)) pending_tasks.append(task) @@ -374,11 +373,11 @@ async def _build_rollout_states(self, rollout_state: RolloutState, item: AgentRo data = await trace_store.export_training_trace.remote(str(rollout_state.session_id), prompt_text) segment_state.input_ids = data["input_ids"] segment_state.labels = data["labels"] - # Agentic training consumes input_ids/labels directly. response_ids is - # filled here only so rollout throughput logging can print rollout_tgs. - segment_state.response_ids = [ - token_id for token_id, label in zip(data["input_ids"][1:], data["labels"][1:]) if label != -100 - ] + # Unified response_ids convention: the contiguous suffix of input_ids after the prompt + # (env-injected tokens included), aligned with response_model_steps for per-token + # staleness; also the token count used by rollout throughput logging. + prompt_len = len(segment_state.prompt_ids or []) + segment_state.response_ids = list(data["input_ids"][prompt_len:]) segment_state.logprobs = data["logprobs"] # ``routed_experts`` is a per-node list for a MoE trace or ``None`` for a dense one; trace_store's # ``_resolve_routed_experts`` already enforces the all-or-nothing invariant (a trace mixing MoE turns with diff --git a/xtuner/v1/rl/trainer/controller.py b/xtuner/v1/rl/trainer/controller.py index 14813351dc..240e128acc 100644 --- a/xtuner/v1/rl/trainer/controller.py +++ b/xtuner/v1/rl/trainer/controller.py @@ -10,7 +10,7 @@ import torch from typing_extensions import NotRequired -from xtuner.v1.data_proto.rl_data import AGENTIC_AGENT_LOOP_TYPES, RolloutState, is_valid_for_training +from xtuner.v1.data_proto.rl_data import RolloutState, is_valid_for_training from xtuner.v1.data_proto.sequence_context import SequenceContext from xtuner.v1.model.compose.base import BaseComposeConfig from xtuner.v1.rl.distillation import DistillationTrainerAdapter @@ -50,21 +50,30 @@ class PackedBatch(TypedDict): teacher_indices: torch.Tensor | None +class _SampleObservation(TypedDict): + """Per-sample metrics observed while converting one rollout state.""" + + advantage: float + supervised_positions: int + prompt_len: int + tool_turns: int | None + training_tokens: int + + def get_train_seq_ctx( input_ids: torch.LongTensor, position_ids: np.ndarray | None = None, multimodal_train_info: dict | None = None, - len_response_ids: int = 0, ) -> SequenceContext: """Build a CPU ``SequenceContext`` for one training sample. Args: input_ids (torch.LongTensor): Model input tokens with shape ``(1, seq_len)``. position_ids (np.ndarray | None): Optional position ids. A 3D array triggers - the VLM MRoPE layout and the response segment is appended. + the VLM MRoPE layout; a prompt-only position segment is extended to cover + the whole input. multimodal_train_info (dict | None): Optional multimodal payload with ``pixel_values``, ``image_grid_thw`` and ``num_img_tokens``. - len_response_ids (int): Response segment length used to extend 3D position ids. Returns: SequenceContext: The CPU sequence context of the sample. @@ -74,22 +83,19 @@ def get_train_seq_ctx( if position_ids is not None and len(position_ids.shape) == 3: # Match get_rope_index_3: response text continues from a single global # max(T, H, W), not per-axis maxima. Per-axis max diverges when the - # prompt ends on image tokens (T≈0 while H/W are large). - max_value = position_ids.amax() - response_position_ids = ( - ( - torch.arange( - 1, - len_response_ids + 1, - device=position_ids.device, - dtype=position_ids.dtype, - ) - + max_value + # prompt ends on image tokens (T≈0 while H/W are large). The extension + # length is derived from the position deficit so samples whose positions + # already cover the whole input (full-sequence agentic samples) need no + # extra shape information. + num_extend = input_ids.size(-1) - position_ids.size(-1) + if num_extend > 0: + max_value = position_ids.amax() + response_position_ids = ( + (torch.arange(1, num_extend + 1, device=position_ids.device, dtype=position_ids.dtype) + max_value) + .view(1, 1, -1) + .expand(3, 1, -1) ) - .view(1, 1, -1) - .expand(3, 1, -1) - ) - position_ids = torch.cat([position_ids, response_position_ids], dim=-1) + position_ids = torch.cat([position_ids, response_position_ids], dim=-1) seq_ctx.position_ids = position_ids # type: ignore[assignment] assert position_ids.size(-1) == input_ids.size(-1) @@ -346,7 +352,7 @@ def _grouped_by_max_length(self, packed_data_batches): # 排序后这条 pack 会被放在最前面,导致 rank0 的第一个 step 消耗的有效 token 数往往少于其他 rank,是正常现象。 return sorted(packed_data_batches, key=lambda x: x["seq_ctx"].max_length_q, reverse=True) - def _convert_rollout_groups( + def _rollout_groups_to_colate_items( self, rollout_groups: list[list[RolloutState]], pack_max_length: int, @@ -358,9 +364,10 @@ def _convert_rollout_groups( Samples arrive canonicalized from the agent loop (``input_ids``/``labels``/ ``logprobs`` share one full-sequence length) with token staleness already baked - into ``labels``. This method shifts by one position, builds tensors and sequence - contexts, computes group-level advantages from session clustered rewards, and - summarizes the ``data_info`` metrics. + into ``labels``. The conversion runs in phases: (1) cluster group rewards by + session, (2) estimate group-level advantages, (3) shift and tensorize each + rollout state into a ``ColateItem``, and (4) summarize the ``data_info`` + metrics. Args: rollout_groups (list[list[RolloutState]]): Rollout groups selected for training. @@ -378,137 +385,223 @@ def _convert_rollout_groups( for state in group ) - # Per-session rewards for distribution metrics. Agentic sessions may split into several - # trainable segments that share one reward; counting that reward once per session keeps - # rewards/* from being weighted by segment count. cluster_rewards_list: list[float] = [] distillation_reward_observations: list[tuple[RolloutState, float]] = [] - advantages_list: list[float] = [] - prompt_len_list: list[float] = [] - response_len_list: list[float] = [] - tool_turns_list: list[int] = [] - training_tokens = 0 + observations: list[_SampleObservation] = [] training_samples = 0 data_batches: list[ColateItem] = [] - for group in rollout_groups: if not is_valid_for_training(group, self.logger): self.logger.error(f"Skip one data group {group} due to rollout failed or empty response.") continue training_samples += len(group) - # Collect rewards independently from task-advantage computation. Pure OPD may omit - # rewards entirely; when rewards are present they remain useful observability signals. - # session_id 只由 agentic loop / XTUNER_DETERMINISTIC 写入;普通 RL 回退到 rollout_id。 - rewards_by_session: dict[Any, float] = {} - session_representatives: list[RolloutState] = [] - cluster_indices: list[int | None] = [] - for state in group: - # 有可能有重复,但是没有其他更好办法 - turns = state.extra_fields.get("agent_tool_turns") - if isinstance(turns, int): - tool_turns_list.append(turns) - if state.reward is None or "score" not in state.reward: - if self.task_adv_weight > 0: - raise ValueError( - f"Reward is missing or does not contain 'score' key in data: {state}, " - f"but task_adv_weight={self.task_adv_weight} > 0" - ) - cluster_indices.append(None) - continue - reward = float(state.reward["score"]) - session_key = state.session_id if state.session_id is not None else state.rollout_id - if session_key not in rewards_by_session: - rewards_by_session[session_key] = reward - session_representatives.append(state) - cluster_indices.append(len(session_representatives) - 1) - - cluster_rewards = list(rewards_by_session.values()) + # Phase 1: session reward clustering. Collect rewards independently from + # task-advantage computation: pure OPD may omit rewards entirely; when rewards + # are present they remain useful observability signals. + cluster_rewards, session_representatives, cluster_indices = self._cluster_group_rewards(group) cluster_rewards_list.extend(cluster_rewards) distillation_reward_observations.extend(zip(session_representatives, cluster_rewards)) - if self.task_adv_weight != 0 and self.advantage_estimator is not None: - cluster_advantages = self.advantage_estimator.compute( - torch.tensor(cluster_rewards, dtype=torch.float32), session_representatives - ) - sample_advantages = [ - float(cluster_advantages[index].item()) if index is not None else 0.0 for index in cluster_indices - ] - else: - sample_advantages = [0.0] * len(group) + # Phase 2: group-level advantage estimation. + sample_advantages = self._compute_group_advantages( + cluster_rewards, session_representatives, cluster_indices + ) + # Phase 3: per-sample shift, tensorization and seq_ctx construction. for i, state in enumerate(group): - input_ids = state.input_ids - labels = state.labels - assert input_ids is not None and labels is not None and len(input_ids) == len(labels), ( - f"Rollout state is not canonicalized for training: {state}" + data_batch, observation = self._convert_rollout_state(state, sample_advantages[i], has_3d_position) + assert observation["training_tokens"] <= pack_max_length, ( + f"{observation['training_tokens']} vs {pack_max_length}" ) - input_len = len(input_ids) - shifted_labels = labels[1:] - - input_ids_t = cast(torch.LongTensor, torch.tensor(input_ids[:-1], dtype=torch.int64).unsqueeze(0)) - shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) - - rollout_logprobs: torch.Tensor | None = None - if state.logprobs is not None: - assert len(state.logprobs) == input_len, f"{len(state.logprobs)} vs {input_len}, data: {state}" - rollout_logprobs = torch.tensor(state.logprobs[1:], dtype=torch.float32).unsqueeze(0) - assert rollout_logprobs.size() == shifted_labels_t.size(), ( - f"{rollout_logprobs.size()} vs {shifted_labels_t.size()}" + data_batches.append(data_batch) + observations.append(observation) + + if not XTUNER_DETERMINISTIC: + random.shuffle(data_batches) + + # Phase 4: step statistics summarization. + info_dict = self._summarize_data_info( + observations, + cluster_rewards_list, + distillation_reward_observations, + training_samples, + raw_rewards_sum=raw_rewards_sum, + raw_rewards_count=raw_rewards_count, + ) + return data_batches, info_dict + + def _cluster_group_rewards( + self, group: list[RolloutState] + ) -> tuple[list[float], list[RolloutState], list[int | None]]: + # 按 session 聚类奖励:agentic session 可能拆成多个共享同一 reward 的可训练 segment, + # 每个 session 只计一次,避免 rewards/* 与 advantage 被 segment 数放大。 + # session_id 只由 agentic loop / XTUNER_DETERMINISTIC 写入;普通 RL 回退到 rollout_id。 + rewards_by_session: dict[Any, float] = {} + session_representatives: list[RolloutState] = [] + cluster_indices: list[int | None] = [] + for state in group: + if state.reward is None or "score" not in state.reward: + if self.task_adv_weight > 0: + raise ValueError( + f"Reward is missing or does not contain 'score' key in data: {state}, " + f"but task_adv_weight={self.task_adv_weight} > 0" ) + cluster_indices.append(None) + continue + reward = float(state.reward["score"]) + session_key = state.session_id if state.session_id is not None else state.rollout_id + if session_key not in rewards_by_session: + rewards_by_session[session_key] = reward + session_representatives.append(state) + cluster_indices.append(len(session_representatives) - 1) + return list(rewards_by_session.values()), session_representatives, cluster_indices + + def _compute_group_advantages( + self, + cluster_rewards: list[float], + session_representatives: list[RolloutState], + cluster_indices: list[int | None], + ) -> list[float]: + """Scatter cluster-level advantage estimates back to group members. - # Keep the advantage layout aligned with input_ids (response excludes EOS). - # Prompt positions predict prompt tokens whose shifted labels are -100, so their - # advantage is 0; the last entry of response_ids (EOS) is only a label, never an input. - advantage_val = sample_advantages[i] - actual_advantages = [0.0 if label == -100 else advantage_val for label in shifted_labels] - advantages_list.extend(advantage_val for label in shifted_labels if label != -100) - - assert input_len - 1 <= pack_max_length, f"{input_len - 1} vs {pack_max_length}" - training_tokens += input_len - 1 - - # prompt+response(reasoning)样本按 prompt/response 原始长度统计分母; - # agentic 样本按 -100 占位统计。 - if state.agent_loop_type in AGENTIC_AGENT_LOOP_TYPES or state.response_ids is None: - prompt_len = sum(label == -100 for label in shifted_labels) - response_len = len(shifted_labels) - prompt_len - else: - response_len = len(state.response_ids) - prompt_len = input_len - response_len - prompt_len_list.append(prompt_len) - response_len_list.append(response_len) - - position_ids = state.position_ids - if has_3d_position and (position_ids is None or position_ids.ndim != 3): - seq_len = input_ids_t.size(-1) - text_position_ids = np.arange(seq_len, dtype=np.int64).reshape(1, 1, -1) - # Mixed batches use three-axis MRoPE position IDs. Repeat the normal text - # positions on every axis so text samples have the same rank as the VLM - # samples when their sequence contexts are packed together. - position_ids = np.broadcast_to(text_position_ids, (3, 1, seq_len)).copy() - multimodal_train_info = cast(dict | None, state.mm_info) - len_response_ids = ( - len(state.response_ids) - 1 - if state.response_ids is not None and state.agent_loop_type not in AGENTIC_AGENT_LOOP_TYPES - else 0 - ) - seq_ctx = get_train_seq_ctx(input_ids_t, position_ids, multimodal_train_info, len_response_ids) - seq_ctx.rollout_routed_experts = state.routed_experts # type: ignore[assignment] # n,layer*expert + Args: + cluster_rewards (list[float]): Session-clustered rewards of the group. + session_representatives (list[RolloutState]): One representative state per session cluster. + cluster_indices (list[int | None]): Per-state cluster index; ``None`` marks reward-missing states. - teacher_fields = self.distillation.rollout_teacher_targets(state, shifted_labels=shifted_labels) + Returns: + list[float]: Advantage per state, aligned with the group order. Zero when task advantage + is disabled or the state carries no reward. + """ + if self.task_adv_weight != 0 and self.advantage_estimator is not None: + cluster_advantages = self.advantage_estimator.compute( + torch.tensor(cluster_rewards, dtype=torch.float32), session_representatives + ) + return [float(cluster_advantages[index].item()) if index is not None else 0.0 for index in cluster_indices] + return [0.0] * len(cluster_indices) - data_dict: ColateItem = { - "seq_ctx": seq_ctx, - "shifted_labels": shifted_labels_t, - "advantage": actual_advantages, - "rollout_logprobs": rollout_logprobs, - } - data_dict.update(cast(ColateItem, teacher_fields)) - data_batches.append(data_dict) + def _convert_rollout_state( + self, + state: RolloutState, + advantage_val: float, + has_3d_position: bool, + ) -> tuple[ColateItem, _SampleObservation]: + """Shift, tensorize and build the sequence context of one rollout + state. - if not XTUNER_DETERMINISTIC: - random.shuffle(data_batches) + Args: + state (RolloutState): Canonicalized rollout state with full-sequence ``input_ids``/``labels``. + advantage_val (float): Group-level advantage broadcast to every supervised position. + has_3d_position (bool): Whether any sample in the batch carries 3D MRoPE position ids. + + Returns: + tuple[ColateItem, _SampleObservation]: The packable training item and the per-sample + metrics observed during conversion. + """ + input_ids = state.input_ids + labels = state.labels + assert input_ids is not None and labels is not None and len(input_ids) == len(labels), ( + f"Rollout state is not canonicalized for training: {state}" + ) + input_len = len(input_ids) + shifted_labels = labels[1:] + + input_ids_t = cast(torch.LongTensor, torch.tensor(input_ids[:-1], dtype=torch.int64).unsqueeze(0)) + shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) + + rollout_logprobs: torch.Tensor | None = None + if state.logprobs is not None: + assert len(state.logprobs) == input_len, f"{len(state.logprobs)} vs {input_len}, data: {state}" + rollout_logprobs = torch.tensor(state.logprobs[1:], dtype=torch.float32).unsqueeze(0) + assert rollout_logprobs.size() == shifted_labels_t.size(), ( + f"{rollout_logprobs.size()} vs {shifted_labels_t.size()}" + ) + + # Keep the advantage layout aligned with input_ids (response excludes EOS). + # Prompt positions predict prompt tokens whose shifted labels are -100, so their + # advantage is 0; the last entry of response_ids (EOS) is only a label, never an input. + actual_advantages = [0.0 if label == -100 else advantage_val for label in shifted_labels] + + # prompt_len 统计原始输入 prompt 长度(VLM 的 prompt 在 train_prompt_ids); + # response_len 统计 LLM 生成的 token 数(labels 非 -100 的监督位)。环境/工具 + # 插入的语义洞不计入两者;prompt_ids 缺失时 prompt_len 退回洞计数。 + prompt_ids = state.extra_fields.get("train_prompt_ids") or state.prompt_ids + prompt_len = len(prompt_ids) if prompt_ids else sum(label == -100 for label in shifted_labels) + response_len = sum(label != -100 for label in shifted_labels) + + position_ids = state.position_ids + if has_3d_position and (position_ids is None or position_ids.ndim != 3): + seq_len = input_ids_t.size(-1) + text_position_ids = np.arange(seq_len, dtype=np.int64).reshape(1, 1, -1) + # Mixed batches use three-axis MRoPE position IDs. Repeat the normal text + # positions on every axis so text samples have the same rank as the VLM + # samples when their sequence contexts are packed together. + position_ids = np.broadcast_to(text_position_ids, (3, 1, seq_len)).copy() + multimodal_train_info = cast(dict | None, state.mm_info) + seq_ctx = get_train_seq_ctx(input_ids_t, position_ids, multimodal_train_info) + seq_ctx.rollout_routed_experts = state.routed_experts # type: ignore[assignment] # n,layer*expert + + teacher_fields = self.distillation.rollout_teacher_targets(state, shifted_labels=shifted_labels) + + data_dict: ColateItem = { + "seq_ctx": seq_ctx, + "shifted_labels": shifted_labels_t, + "advantage": actual_advantages, + "rollout_logprobs": rollout_logprobs, + } + data_dict.update(cast(ColateItem, teacher_fields)) + + # 有可能有重复,但是没有其他更好办法 + turns = state.extra_fields.get("agent_tool_turns") + observation: _SampleObservation = { + "advantage": advantage_val, + "supervised_positions": response_len, + "prompt_len": prompt_len, + "tool_turns": turns if isinstance(turns, int) else None, + "training_tokens": input_len - 1, + } + return data_dict, observation + + def _summarize_data_info( + self, + observations: list[_SampleObservation], + cluster_rewards_list: list[float], + distillation_reward_observations: list[tuple[RolloutState, float]], + training_samples: int, + raw_rewards_sum: float, + raw_rewards_count: int, + ) -> dict[str, float]: + """Reduce phase-1/phase-3 observations into the step ``data_info`` + metrics. + + Args: + observations (list[_SampleObservation]): Per-sample metrics collected in phase 3. + cluster_rewards_list (list[float]): Session-clustered rewards collected in phase 1. + distillation_reward_observations (list[tuple[RolloutState, float]]): Session + representative states paired with their cluster reward. + training_samples (int): Number of states in the valid groups selected for training. + raw_rewards_sum (float): Producer-side raw reward sum used by ``raw_rewards/mean``. + raw_rewards_count (int): Producer-side raw reward count used by ``raw_rewards/mean``. + + Returns: + dict[str, float]: The step statistics dictionary. + """ + advantages_list: list[float] = [] + prompt_len_list: list[float] = [] + response_len_list: list[float] = [] + tool_turns_list: list[int] = [] + training_tokens = 0 + for observation in observations: + # Advantage metrics count one entry per supervised position, matching the + # per-position advantage layout built in phase 3. + advantages_list.extend([observation["advantage"]] * observation["supervised_positions"]) + prompt_len_list.append(observation["prompt_len"]) + response_len_list.append(observation["supervised_positions"]) + if observation["tool_turns"] is not None: + tool_turns_list.append(observation["tool_turns"]) + training_tokens += observation["training_tokens"] # rewards/* report the per-session reward distribution; batch_size/training_samples # count the valid rollout segments included in training. @@ -543,7 +636,7 @@ def _convert_rollout_groups( info_dict["tool_turns/mean"] = tool_turns_t.mean().item() info_dict["tool_turns/min"] = float(tool_turns_t.min().item()) info_dict["tool_turns/max"] = float(tool_turns_t.max().item()) - return data_batches, info_dict + return info_dict def fit( self, @@ -553,7 +646,7 @@ def fit( raw_rewards_sum: float = 0.0, raw_rewards_count: int = 0, ) -> tuple[list[WorkerLogItem], dict[str, float]]: - data_batches, data_info = self._convert_rollout_groups( + data_batches, data_info = self._rollout_groups_to_colate_items( rollout_groups, pack_max_length, raw_rewards_sum=raw_rewards_sum, From d92d059b1fee8df0258d77dd7cfdab2665140695 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Wed, 23 Sep 2026 03:26:05 +0000 Subject: [PATCH 3/4] [Test] Align tests with labels-based staleness and sequence conventions --- tests/rl/test_producer.py | 3 ++ .../test_qwen35_vl_moe_async_train_2step.py | 5 ++- tests/rl/test_replay_buffer.py | 38 +++++++++++++------ tests/rl/test_rl_trainer_checkpoint.py | 2 +- tests/rl/test_trajectory_logging.py | 1 + 5 files changed, 35 insertions(+), 14 deletions(-) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index c4933e0f15..6c7c493a23 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -278,6 +278,7 @@ async def test_disagg_put_uses_consumer_step_for_token_expiry(self): tokens=[1, 11], response="old response", response_ids=[11], + labels=[1], response_model_steps=[3], logprobs=[0.1], finish_reason="stop", @@ -758,12 +759,14 @@ async def test_async_produce_strategy_rerolls_expired_state_and_preserves_fresh_ expired = make_rollout_state(900) expired.response = "expired response" expired.response_ids = [11] + expired.labels = [1] expired.response_model_steps = [0] expired.reward = {"score": 0.1} fresh = make_rollout_state(901) fresh.group_id = expired.group_id fresh.response = "fresh response" fresh.response_ids = [21] + fresh.labels = [1] fresh.response_model_steps = [5] fresh.logprobs = [-0.2] fresh.reward = {"score": 0.9} diff --git a/tests/rl/test_qwen35_vl_moe_async_train_2step.py b/tests/rl/test_qwen35_vl_moe_async_train_2step.py index 8036152869..433b6af5e3 100644 --- a/tests/rl/test_qwen35_vl_moe_async_train_2step.py +++ b/tests/rl/test_qwen35_vl_moe_async_train_2step.py @@ -443,7 +443,10 @@ def _assert_vlm_rollout_states(self) -> None: self.assertIn("image_data", sample.extra_fields) self.assertTrue(sample.response_ids) self.assertIsNotNone(sample.logprobs) - self.assertEqual(len(sample.logprobs), len(sample.response_ids)) + # canonicalize 后 input_ids/labels/logprobs 全序列对齐(prompt 前缀补 0)。 + expected_total_len = len(sample.extra_fields["train_prompt_ids"]) + len(sample.response_ids) + self.assertEqual(len(sample.input_ids), expected_total_len) + self.assertEqual(len(sample.logprobs), len(sample.input_ids)) self.assertIsNotNone(sample.reward) self.assertIn("score", sample.reward) diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index 691f125a97..7c8c117424 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -291,6 +291,7 @@ async def test_common_put_token_expiry_preserves_fresh_group_members(self): 1, response="expired response", response_ids=[11, 12], + labels=[1, 1], response_model_steps=[0, 0], reward={"score": 0.1}, ) @@ -298,6 +299,7 @@ async def test_common_put_token_expiry_preserves_fresh_group_members(self): 2, response="fresh response", response_ids=[21, 22], + labels=[1, 1], response_model_steps=[4, 4], reward={"score": 0.9}, ) @@ -323,20 +325,28 @@ async def test_common_put_token_expiry_preserves_fresh_group_members(self): self.assertEqual(group[1].response_model_steps, [4, 4]) self.assertEqual(group[1].reward, {"score": 0.9}) - async def test_common_put_skips_token_expiry_for_agentic_group(self): + async def test_common_put_token_expiry_applies_to_agentic_group(self): + # labels 对齐的 agentic state 不再跳过 token expiry:过期成员整体过期,新鲜成员原样保留。 for config_name, replay_buffer_config_cls in REPLAY_BUFFER_CONFIGS: with self.subTest(replay_buffer_config=config_name): replay_buffer = replay_buffer_config_cls().build() - state = make_rollout_state( + stale = make_rollout_state( 1, - response="agentic response", + response="stale agentic response", response_model_steps=[0], input_ids=[1, 2], labels=[-100, 2], ) + fresh = make_rollout_state( + 2, + response="fresh agentic response", + response_model_steps=[4], + input_ids=[3, 4], + labels=[-100, 4], + ) await replay_buffer.put( - [state], + [stale, fresh], "task", current_train_step=5, stale_threshold=10, @@ -344,10 +354,14 @@ async def test_common_put_skips_token_expiry_for_agentic_group(self): expired_groups_retryable=True, ) - self.assertEqual(await replay_buffer.count("task", Status.COMPLETED), 1) - self.assertEqual(await replay_buffer.count("task", Status.EXPIRED), 0) - self.assertEqual(state.status, Status.COMPLETED) - self.assertEqual(state.response, "agentic response") + self.assertEqual(await replay_buffer.count("task", Status.COMPLETED), 0) + self.assertEqual(await replay_buffer.count("task", Status.EXPIRED), 1) + group = (await replay_buffer.get(1, "task", Status.EXPIRED))[0] + self.assertEqual([item.status for item in group], [Status.EXPIRED, Status.COMPLETED]) + self.assertEqual(group[0].response, "") + self.assertEqual(group[1].response, "fresh agentic response") + self.assertEqual(group[1].input_ids, [3, 4]) + self.assertEqual(group[1].labels, [-100, 4]) async def test_common_put_seq_expiry_preserves_fresh_group_members(self): # seq expiry 路由整组到 EXPIRED pool,但只标记和清理实际过期的 state。 @@ -382,8 +396,8 @@ async def test_common_put_drops_entire_token_expired_group_when_rerollout_is_dis for config_name, replay_buffer_config_cls in REPLAY_BUFFER_CONFIGS: with self.subTest(replay_buffer_config=config_name): replay_buffer = replay_buffer_config_cls().build() - expired = make_rollout_state(1, response_model_steps=[0]) - fresh = make_rollout_state(2, response_model_steps=[4]) + expired = make_rollout_state(1, labels=[1], response_model_steps=[0]) + fresh = make_rollout_state(2, labels=[1], response_model_steps=[4]) await replay_buffer.put( [expired, fresh], @@ -405,8 +419,8 @@ async def test_common_refresh_token_expiry_moves_mixed_group_to_expired_pool(sel for config_name, replay_buffer_config_cls in REPLAY_BUFFER_CONFIGS: with self.subTest(replay_buffer_config=config_name): replay_buffer = replay_buffer_config_cls().build() - expired = make_rollout_state(1, response_model_steps=[0], reward={"score": 0.1}) - fresh = make_rollout_state(2, response_model_steps=[4], reward={"score": 0.9}) + expired = make_rollout_state(1, labels=[1], response_model_steps=[0], reward={"score": 0.1}) + fresh = make_rollout_state(2, labels=[1], response_model_steps=[4], reward={"score": 0.9}) await replay_buffer.put([expired, fresh], "task") expired_counts = await replay_buffer.refresh_staleness( diff --git a/tests/rl/test_rl_trainer_checkpoint.py b/tests/rl/test_rl_trainer_checkpoint.py index e139c7802b..3f132ddbaa 100644 --- a/tests/rl/test_rl_trainer_checkpoint.py +++ b/tests/rl/test_rl_trainer_checkpoint.py @@ -211,7 +211,7 @@ def build_pg(resources_config, name="train"): build_pg.built = getattr(build_pg, "built", []) + [name] return SimpleNamespace(id=f"pg-{idx}", bundle_specs=[]) - def build_train_controller(worker_cfg, placement_group): + def build_train_controller(worker_cfg, placement_group, advantage_estimator=None, distillation=None): controller = _FakeTrainController() runtime.train_controllers.append(controller) return controller diff --git a/tests/rl/test_trajectory_logging.py b/tests/rl/test_trajectory_logging.py index 14f7ebd194..38c0d90e8f 100644 --- a/tests/rl/test_trajectory_logging.py +++ b/tests/rl/test_trajectory_logging.py @@ -32,6 +32,7 @@ def _make_state() -> RolloutState: rollout_id=7, session_id=12345678901234567890, message=[{"role": "user", "content": "test"}], + prompt_ids=[1], num_tokens=1, extra_fields={}, ) From 92ab26eb0b116bcf25bfbc8a97f368108d387467 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Wed, 23 Sep 2026 06:19:12 +0000 Subject: [PATCH 4/4] [Fix] Apply review fixes: pure staleness masks, agentic logprobs padding, multimodal assert --- .../verl_agent/common/agent_loop_verl_tool.py | 4 +- tests/rl/test_staleness_policy.py | 48 ++++++++++++++----- xtuner/v1/data_proto/rl_data.py | 29 +++++------ xtuner/v1/rl/agent_loop/gsm8k_with_tool.py | 2 +- .../agent_in_localhost_loop.py | 6 +++ .../agent_in_sandbox_loop.py | 6 +++ .../v1/rl/agent_loop_manager/produce_utils.py | 10 +++- xtuner/v1/rl/replay_buffer.py | 2 +- 8 files changed, 73 insertions(+), 34 deletions(-) diff --git a/recipe/verl_agent/common/agent_loop_verl_tool.py b/recipe/verl_agent/common/agent_loop_verl_tool.py index 759c5bce95..c30d7a85c5 100644 --- a/recipe/verl_agent/common/agent_loop_verl_tool.py +++ b/recipe/verl_agent/common/agent_loop_verl_tool.py @@ -140,7 +140,9 @@ async def generate_sample(self, rollout_state: RolloutState) -> RolloutState: semantic_mask = output.response_mask if output.response_mask is not None else [1] * len(response_ids) rollout_state.prompt_ids = prompt_ids rollout_state.response_ids = response_ids - rollout_state.logprobs = output.response_logprobs + rollout_state.logprobs = ( + None if output.response_logprobs is None else [0.0] * len(prompt_ids) + list(output.response_logprobs) + ) rollout_state.routed_experts = output.routed_experts rollout_state.input_ids = prompt_ids + response_ids rollout_state.labels = [-100] * len(prompt_ids) + [ diff --git a/tests/rl/test_staleness_policy.py b/tests/rl/test_staleness_policy.py index 8a8785b864..5764aac2ef 100644 --- a/tests/rl/test_staleness_policy.py +++ b/tests/rl/test_staleness_policy.py @@ -83,7 +83,7 @@ def test_token_staleness_threshold_can_be_relaxed(self): with self.subTest(token_stale_threshold=token_stale_threshold): state = self._state(response_model_steps=[0, 4]) - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [state], current_train_step=5, token_stale_threshold=token_stale_threshold, @@ -94,14 +94,14 @@ def test_token_staleness_threshold_can_be_relaxed(self): self.assertEqual(state.labels, expected_labels) def test_repeated_calls_converge(self): - # staleness 只增不减:先在旧 step 烙一次,再在新 step 重算,最终 labels - # 与一次性按新 step 烙制的结果完全一致(replay buffer 逐轮检查同理)。 + # staleness 只增不减:先在旧 step 计算并烙制一次,再在新 step 重算重烙,最终 + # labels 与一次性按新 step 计算并烙制的结果完全一致。 incremental = self._state(response_model_steps=[0, 4]) - calculate_group_effective_response_masks([incremental], current_train_step=5, token_stale_threshold=4) - calculate_group_effective_response_masks([incremental], current_train_step=9, token_stale_threshold=4) + self._calc_and_bake([incremental], current_train_step=5, token_stale_threshold=4) + self._calc_and_bake([incremental], current_train_step=9, token_stale_threshold=4) oneshot = self._state(response_model_steps=[0, 4]) - calculate_group_effective_response_masks([oneshot], current_train_step=9, token_stale_threshold=4) + self._calc_and_bake([oneshot], current_train_step=9, token_stale_threshold=4) self.assertEqual(incremental.labels, oneshot.labels) self.assertEqual(oneshot.labels, [-100, -100, -100, -100]) @@ -110,7 +110,7 @@ def test_token_staleness_intersects_semantic_labels(self): # 最终有效掩码必须同时反映语义监督(labels 尾段 -100 位)与 token staleness。 state = self._state(response_model_steps=[0, 4], labels=[-100, -100, 3, -100]) - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [state], current_train_step=5, token_stale_threshold=4, @@ -127,7 +127,7 @@ def test_rerolled_state_uses_token_staleness_only(self): state.response_model_steps = [4, 4] state.labels = [-100, -100, 3, 4] - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [state], current_train_step=5, token_stale_threshold=4, @@ -161,7 +161,7 @@ def test_full_sequence_agentic_stale_clears_all_supervised_labels(self): labels=[-100, -100, 30, -100, 40, -100, 32], ) - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [state], current_train_step=5, token_stale_threshold=4, @@ -184,7 +184,7 @@ def test_full_sequence_agentic_per_token_stale_clears_only_stale_tokens(self): labels=[-100, -100, 30, -100, 40, -100, 32], ) - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [state], current_train_step=5, token_stale_threshold=4, @@ -206,7 +206,7 @@ def test_full_sequence_agentic_fresh_keeps_labels(self): labels=[-100, -100, 30, -100, 40, -100, 32], ) - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [state], current_train_step=5, token_stale_threshold=4, @@ -227,7 +227,7 @@ def test_zero_prompt_suffix_state_is_eligible(self): labels=[3, 4], ) - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [state], current_train_step=5, token_stale_threshold=4, @@ -242,7 +242,7 @@ def test_canonicalized_prompt_response_state_is_still_eligible(self): canonical.input_ids = [1, 2, 3, 4] canonical.logprobs = [0.0, 0.0, -0.1, -0.2] - masks = calculate_group_effective_response_masks( + masks = self._calc_and_bake( [canonical], current_train_step=5, token_stale_threshold=4, @@ -251,6 +251,28 @@ def test_canonicalized_prompt_response_state_is_still_eligible(self): self.assertEqual(masks, [[0, 0]]) self.assertEqual(canonical.labels, [-100, -100, -100, -100]) + @staticmethod + def _calc_and_bake( + group: list[RolloutState], + *, + current_train_step: int, + token_stale_threshold: int, + ) -> list[list[int] | None]: + masks = calculate_group_effective_response_masks( + group, + current_train_step=current_train_step, + token_stale_threshold=token_stale_threshold, + ) + for item, mask in zip(group, masks): + if mask is None: + continue + offset = len(item.labels) - len(mask) + item.labels[offset:] = [ + label if mask_value else -100 + for label, mask_value in zip(item.labels[offset:], mask) + ] + return masks + @staticmethod def _state( *, diff --git a/xtuner/v1/data_proto/rl_data.py b/xtuner/v1/data_proto/rl_data.py index 512cf6ce71..1173f31e88 100644 --- a/xtuner/v1/data_proto/rl_data.py +++ b/xtuner/v1/data_proto/rl_data.py @@ -416,18 +416,18 @@ def _calculate_effective_response_mask( current_train_step: int, token_stale_threshold: int, ) -> list[int]: - """Bake token staleness into labels and return the effective response mask. + """Calculate the effective response mask of one rollout state without + mutating it. Every loop writes ``response_ids`` under one convention: the contiguous suffix of ``input_ids`` after the prompt (env-injected or tool tokens included), so response token ``j`` maps to labels position ``len(labels) - len(response_ids) + j``. The effective mask follows the reasoning-RL rule ``effective = semantic_mask * token_staleness_mask``: the semantic mask is recovered from ``labels != -100`` on the response region (semantic holes stay supervised-out regardless of staleness), - staleness is evaluated per token via ``response_model_steps``, and a zero effective mask bakes ``-100`` - into the label in place. + and staleness is evaluated per token via ``response_model_steps``. Args: - rollout_state (RolloutState): Rollout sample whose labels are updated in place. + rollout_state (RolloutState): Rollout sample to inspect; left unmodified. current_train_step (int): Trainer step that will consume the sample. token_stale_threshold (int): Maximum token staleness, measured in trainer steps, allowed for training. @@ -451,9 +451,6 @@ def _calculate_effective_response_mask( semantic_mask_value * token_staleness_mask_value for semantic_mask_value, token_staleness_mask_value in zip(semantic_mask, token_staleness_mask) ] - for i, mask_value in enumerate(effective_mask): - if mask_value == 0: - labels[offset + i] = -100 return effective_mask @@ -463,18 +460,18 @@ def calculate_group_effective_response_masks( current_train_step: int, token_stale_threshold: int | None, ) -> list[list[int] | None]: - """Calculate a group's effective masks and bake token staleness into its - labels. + """Calculate one group's effective response masks under the token-staleness + policy. - For each eligible state, stale supervised labels are cleared to ``-100`` in place. - Clearing only ever extends: staleness grows monotonically with the trainer step, so - repeated calls (e.g. replay-buffer expiry checks followed by the train-batch bake) - converge to the same labels. Each returned mask is the effective response mask - aligned with ``response_ids``. ``None`` means token staleness is disabled or does - not apply to that state (no labels, or no ``response_ids`` to align with). + The function is pure: ``group`` is inspected, not modified. Each returned mask is the effective + response mask aligned with ``response_ids``; callers bake them into ``labels`` when the train + batch is taken, clearing zero-mask supervised positions to ``-100``. The clearing only ever + extends, so masks computed at growing trainer steps converge to the same labels. ``None`` means + token staleness is disabled or does not apply to that state (no labels, or no ``response_ids`` to + align with). Args: - group (list[RolloutState]): Rollout group updated in place. + group (list[RolloutState]): Rollout group to inspect; left unmodified. current_train_step (int): Trainer step that will consume the group. token_stale_threshold (int | None): Maximum token staleness, measured in trainer steps, allowed for training. ``None`` disables token staleness. diff --git a/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py b/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py index 4108a48765..bdd7f8c117 100644 --- a/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py +++ b/xtuner/v1/rl/agent_loop/gsm8k_with_tool.py @@ -157,7 +157,7 @@ async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> Rollou f"Prompt ids cannot be None or empty in data: {rollout_state}" ) rollout_state.response_ids = final_response_ids - rollout_state.logprobs = final_logprobs + rollout_state.logprobs = [0.0] * len(prompt_ids) + final_logprobs rollout_state.input_ids = list(prompt_ids) + final_response_ids rollout_state.labels = [-100] * len(prompt_ids) + [ resp_id if mask_value else -100 for resp_id, mask_value in zip(final_response_ids, final_response_mask) diff --git a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py index 00c28e6747..509eca487f 100644 --- a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py +++ b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py @@ -258,6 +258,12 @@ async def _fill_rollout_state(self, rollout_state: RolloutState, item: AgentRoll self._fill_eval_rollout_state(rollout_state, item) return + # prompt_ids 切片只对纯文本 trace 成立;train_prompt_ids(VLM 扩充版 prompt)一旦存在, + # trace input_ids 与 prompt_ids 的长度约定即失效,宁可快速失败也不静默产出错位训练字段。 + assert "train_prompt_ids" not in rollout_state.extra_fields, ( + f"Agentic trace rollout does not support multimodal samples, rollout_id={rollout_state.rollout_id}" + ) + response_message = item.artifacts.get("response_message") or {} rollout_state.status = Status.COMPLETED if item.status == RolloutStatus.COMPLETED else Status.FAILED rollout_state.finish_reason = str( diff --git a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py index 19ea00513b..25a4809d60 100644 --- a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py +++ b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py @@ -322,6 +322,12 @@ async def _build_rollout_states(self, rollout_state: RolloutState, item: AgentRo self._fill_eval_rollout_state(rollout_state, item) return [rollout_state] + # prompt_ids 切片只对纯文本 trace 成立;train_prompt_ids(VLM 扩充版 prompt)一旦存在, + # trace input_ids 与 prompt_ids 的长度约定即失效,宁可快速失败也不静默产出错位训练字段。 + assert "train_prompt_ids" not in rollout_state.extra_fields, ( + f"Agentic trace rollout does not support multimodal samples, rollout_id={rollout_state.rollout_id}" + ) + response_message = _response_message(item.artifacts, required=item.status == RolloutStatus.COMPLETED) rollout_state.status = Status.COMPLETED if item.status == RolloutStatus.COMPLETED else Status.FAILED rollout_state.finish_reason = str( diff --git a/xtuner/v1/rl/agent_loop_manager/produce_utils.py b/xtuner/v1/rl/agent_loop_manager/produce_utils.py index 8214bc581a..810f972cff 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -6,7 +6,7 @@ from dataclasses import dataclass from enum import Enum, auto from pathlib import Path -from typing import TYPE_CHECKING, Any, Awaitable, Callable, Protocol, runtime_checkable +from typing import TYPE_CHECKING, Any, Awaitable, Callable, Protocol, cast, runtime_checkable import ray import tqdm @@ -593,11 +593,17 @@ async def take_train_batch( if task.token_stale_threshold is None: continue for group in batch_by_task.get(task.task_name, []): - calculate_group_effective_response_masks( + effective_masks = calculate_group_effective_response_masks( group, current_train_step=current_train_step, token_stale_threshold=task.token_stale_threshold, ) + for item, mask in zip(group, effective_masks): + if mask is None: + continue + labels = cast(list[int], item.labels) + offset = len(labels) - len(mask) + labels[offset:] = [label if mask_value else -100 for label, mask_value in zip(labels[offset:], mask)] if hasattr(progress, "mark_consumed"): progress.mark_consumed(consumed_counts) diff --git a/xtuner/v1/rl/replay_buffer.py b/xtuner/v1/rl/replay_buffer.py index 4f033ddbb9..1ab1bba61b 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -459,7 +459,7 @@ def _apply_staleness_lifecycle( # 1. update seq-level staleness refresh_seq_staleness(group, current_train_step) - # 2. bake token-level staleness into labels and use the masks for expiry decisions + # 2. Token-level effective masks drive expiry decisions; labels are baked at take_train_batch. token_level_effective_masks = calculate_group_effective_response_masks( group, current_train_step=current_train_step,