From 58528ebf1d79778b606dc87153b38373e5decd31 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Wed, 22 Jul 2026 04:30:52 +0000 Subject: [PATCH 1/8] [Refactor] Move rollout filtering into agent loops --- .../config/reasoning_rl_qwen3p5vl_mtp_ep.py | 4 +-- .../v1/config/rl_dapo_math_async_filter.py | 2 +- tests/rl/test_producer.py | 23 +++++++++----- tests/rl/test_rl_colocate_trainer.py | 2 +- xtuner/v1/rl/agent_loop/agent_loop.py | 29 ++++++++++++++++- .../agent_loop_manager/agent_loop_manager.py | 5 +++ .../disagg_agent_loop_manager.py | 3 ++ .../rl/agent_loop_manager/disagg_producer.py | 11 ++----- .../v1/rl/agent_loop_manager/produce_utils.py | 31 ++++++------------- xtuner/v1/rl/agent_loop_manager/producer.py | 30 ++++-------------- 10 files changed, 73 insertions(+), 67 deletions(-) diff --git a/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py b/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py index 92d038acef..b1bc3a650d 100644 --- a/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py +++ b/examples/v1/config/reasoning_rl_qwen3p5vl_mtp_ep.py @@ -282,10 +282,9 @@ def group_sample_filter_func(group_samples): produce_strategy_config = AsyncProduceStrategyConfig( over_sample_threshold=1, enable_partial_rollout=1, - is_valid_sample_fn=group_sample_filter_func, max_staleness=3, ) -# produce_strategy_config= SyncProduceStrategyConfig(is_valid_sample_fn=group_sample_filter_func) +# produce_strategy_config = SyncProduceStrategyConfig() # 6. agent loop managers agent_loop_config = SingleTurnAgentLoopConfig( @@ -297,6 +296,7 @@ def group_sample_filter_func(group_samples): task_name="train_task", agent_loop_config=agent_loop_config, judger_config=judger_config, + filter_func=group_sample_filter_func, produce_strategy_config=produce_strategy_config, sampler_config=SamplerConfig(dataloader_cfg=dataloader_cfg, prompt_repeat_k=prompt_repeat_k), ), diff --git a/examples/v1/config/rl_dapo_math_async_filter.py b/examples/v1/config/rl_dapo_math_async_filter.py index 6b99b9a1b4..826768c18f 100644 --- a/examples/v1/config/rl_dapo_math_async_filter.py +++ b/examples/v1/config/rl_dapo_math_async_filter.py @@ -157,13 +157,13 @@ def group_samples_filter_func(rollout_states): enable_partial_rollout=True, max_staleness=0, tail_batch_trigger_size=256, - is_valid_sample_fn=group_samples_filter_func ) agent_loop_manager_cfg = AgentLoopManagerConfig( tasks=TaskSpecConfig( task_name="train_task", agent_loop_config=agent_loop_config, judger_config=judger_config, + filter_func=group_samples_filter_func, produce_strategy_config=produce_strategy_config, sampler_config=sampler_config, ), diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 3514542f89..b20815d81a 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -23,6 +23,7 @@ from unittest.mock import AsyncMock, MagicMock from xtuner.v1.data_proto.rl_data import RolloutState, Status, discard_rollout_state +from xtuner.v1.rl.agent_loop import AgentLoop from xtuner.v1.rl.agent_loop_manager import ( AsyncProduceStrategyConfig, DisaggAsyncProduceStrategyConfig, @@ -117,6 +118,7 @@ async def mock_gen(rs, **kwargs): return rs mock_agent_loop.generate_group = mock_gen + mock_agent_loop.collect_rollout_group = AgentLoop.collect_rollout_group.__get__(mock_agent_loop) return mock_agent_loop def _build_context( @@ -130,6 +132,7 @@ def _build_context( train_step: int = 0, model_step: int = 0, progress: ProduceProgress | None = None, + is_valid_sample_fn=None, ) -> ProduceContext: # 测试只走新的 ProduceContext 入口,不再覆盖旧散装参数兼容逻辑。 if progress is None: @@ -143,7 +146,7 @@ def _build_context( train_step=train_step, model_step=model_step, progress=progress, - is_valid_sample_fn=strategy.is_valid_sample_fn, + is_valid_sample_fn=is_valid_sample_fn, stale_threshold=getattr(strategy, "stale_threshold", None), ) @@ -193,7 +196,6 @@ def _build_disagg_context( update_event=update_event, model_step=model_step, progress=progress, - is_valid_sample_fn=strategy.is_valid_sample_fn, stale_threshold=getattr(strategy, "stale_threshold", None), ) @@ -280,8 +282,8 @@ async def test_discard_rollout_state_keeps_required_fields_valid(self): self.assertIsNone(discarded.routed_experts) self.assertEqual(discarded.extra_fields, {}) - async def test_put_generated_group_only_validates_completed_group(self): - # 验证 ProduceContext 只对 completed group 执行业务过滤,aborted group 保持可重试状态。 + async def test_collection_filters_only_completed_group(self): + # 过滤由 AgentLoop collection 执行;Producer 只处理返回状态和数据所有权。 task_name = "test_valid_completed_only" valid_checked_statuses = [] @@ -289,16 +291,18 @@ def is_valid_sample_fn(samples): valid_checked_statuses.append([sample.status for sample in samples]) return False - strategy = SyncProduceStrategyConfig(is_valid_sample_fn=is_valid_sample_fn).build() + strategy = SyncProduceStrategyConfig().build() ctx = self._build_context( strategy, task_name, self._build_agent_loop(), self._build_sampler(), batch_size=1, + is_valid_sample_fn=is_valid_sample_fn, ) completed_group = [make_rollout_state(1, status=Status.COMPLETED)] + completed_group = await ctx.collect_rollout_group(completed_group) self.assertFalse(await ctx.put_generated_group(completed_group)) self.assertIsNone(completed_group[0].uid) @@ -339,19 +343,21 @@ async def test_put_generated_group_records_raw_rewards_before_filtering(self): def is_valid_sample_fn(samples): return False - strategy = SyncProduceStrategyConfig(is_valid_sample_fn=is_valid_sample_fn).build() + strategy = SyncProduceStrategyConfig().build() ctx = self._build_context( strategy, task_name, self._build_agent_loop(), self._build_sampler(), batch_size=1, + is_valid_sample_fn=is_valid_sample_fn, ) completed_group = [ make_rollout_state(1, status=Status.COMPLETED, reward_score=0.25), make_rollout_state(2, status=Status.COMPLETED, reward_score=0.75), ] + completed_group = await ctx.collect_rollout_group(completed_group) self.assertFalse(await ctx.put_generated_group(completed_group)) self.assertTrue(all(item.uid is None for item in completed_group)) @@ -425,7 +431,7 @@ async def mock_gen(rs, **kwargs): mock_agent_loop = self._build_agent_loop() mock_agent_loop.generate_group = mock_gen - strategy = SyncProduceStrategyConfig(is_valid_sample_fn=is_valid_sample_fn).build() + strategy = SyncProduceStrategyConfig().build() sampler = self._build_sampler() ctx = self._build_context( strategy, @@ -436,6 +442,7 @@ async def mock_gen(rs, **kwargs): train_step=4, model_step=3, progress=self._build_progress(task_name, target=2), + is_valid_sample_fn=is_valid_sample_fn, ) await strategy.produce_batch(ctx) @@ -465,7 +472,7 @@ async def mock_gen(rs, **kwargs): r.status = Status.COMPLETED return rs - mock_agent_loop.generate_group = mock_gen + mock_agent_loop.collect_rollout_group = mock_gen sampler_cfg = SamplerConfig.model_construct(dataloader_cfg=self.mock_dataloader_cfg) produce_strategy_cfg = AsyncProduceStrategyConfig(over_sample_threshold=1) diff --git a/tests/rl/test_rl_colocate_trainer.py b/tests/rl/test_rl_colocate_trainer.py index 69f902d6d0..3328ec90c4 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -98,7 +98,7 @@ async def generate_group(rollout_states, **kwargs): state.response_model_steps = [model_step] return rollout_states - agent_loop.generate_group = generate_group + agent_loop.collect_rollout_group = generate_group return agent_loop diff --git a/xtuner/v1/rl/agent_loop/agent_loop.py b/xtuner/v1/rl/agent_loop/agent_loop.py index 6a69d0f80c..22c23186f9 100644 --- a/xtuner/v1/rl/agent_loop/agent_loop.py +++ b/xtuner/v1/rl/agent_loop/agent_loop.py @@ -2,6 +2,7 @@ import asyncio from abc import ABC, abstractmethod +from collections.abc import Callable from typing import Any, TypeAlias, cast, overload import ray @@ -9,7 +10,7 @@ from ray.actor import ActorClass, ActorProxy from ray.util.placement_group import PlacementGroup -from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status, get_group_status from xtuner.v1.rl.judger import Judger from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.rl.rollout.constants import AGENT_LOOP_RAY_GENERATE_MAX_CONCURRENCY @@ -201,6 +202,21 @@ async def generate_group(self, rollout_state: list[RolloutState], **kwargs) -> l group_samples = await self.run_judger(group_samples) return group_samples + async def collect_rollout_group( + self, + rollout_state: list[RolloutState], + *, + is_valid_sample_func: Callable[[list[RolloutState]], bool] | None = None, + **kwargs, + ) -> list[RolloutState]: + group = await self.generate_group(rollout_state, **kwargs) + if get_group_status(group) != Status.COMPLETED: + return group + if is_valid_sample_func is not None and not is_valid_sample_func(group): + for state in group: + state.status = Status.FILTERED + return group + @overload async def run_judger(self, rollout_state: RolloutState) -> RolloutState: ... @@ -287,6 +303,13 @@ async def generate_group(self, rollout_state: list[RolloutState], **kwargs) -> l finally: await self._release_worker(worker) + async def collect_rollout_group(self, rollout_state: list[RolloutState], **kwargs) -> list[RolloutState]: + worker = await self._pick_worker() + try: + return await worker.collect_rollout_group.remote(rollout_state, **kwargs) + finally: + await self._release_worker(worker) + def get_worker_status(self) -> dict[str, int]: return {str(worker): load for worker, load in self._worker_loads.items()} @@ -329,6 +352,10 @@ async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> Rollou async def generate_group(self, rollout_state: list[RolloutState], **kwargs) -> list[RolloutState]: return await self.agent_loop.generate_group(rollout_state, **kwargs) + @ray_method(concurrency_group=AGENT_LOOP_CONCURRENCY_GROUP_GENERATE) + async def collect_rollout_group(self, rollout_state: list[RolloutState], **kwargs) -> list[RolloutState]: + return await self.agent_loop.collect_rollout_group(rollout_state, **kwargs) + @ray_method async def get_rollout_ctl(self): return self.agent_loop.rollout_ctl diff --git a/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py b/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py index 2b45a6ca15..3aa0b0e04c 100644 --- a/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py +++ b/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py @@ -18,6 +18,7 @@ _MANAGER_STATE_PATH, _STATUS_POLL_INTERVAL_S, _TASK_CHECKPOINT_DIR, + IsValidSampleFn, ProduceBatchResult, _TaskRunner, _TaskSamplerView, @@ -55,6 +56,8 @@ class TaskSpecConfig(BaseModel): judger_config (JudgerConfig | ComposedJudgerConfig | None): Optional judger configuration used to score generated samples. Defaults to None. + filter_func (IsValidSampleFn | None): Optional group filter applied by the + agent loop after generation. Defaults to None. produce_strategy_config (ProduceStrategyConfig): Strategy used to produce rollout samples. Defaults to ``SyncProduceStrategyConfig``. sampler_config (SamplerConfig): Dataset sampler configuration for this @@ -82,6 +85,7 @@ class TaskSpecConfig(BaseModel): weight: float = Field(default=1.0, ge=0.0) agent_loop_config: AgentLoopConfig judger_config: JudgerConfig | ComposedJudgerConfig | None = None + filter_func: IsValidSampleFn | None = None produce_strategy_config: ProduceStrategyConfig = SyncProduceStrategyConfig() sampler_config: SamplerConfig @@ -154,6 +158,7 @@ def build( agent_loop=agent_loop, produce_strategy=produce_strategy, sampler=sampler, + is_valid_sample_fn=task_cfg.filter_func, weight=task_cfg.weight, order=order, ) diff --git a/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py b/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py index e73dc63da3..ef2ff03434 100644 --- a/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py +++ b/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py @@ -27,6 +27,7 @@ _MANAGER_STATE_PATH, _STATUS_POLL_INTERVAL_S, _TASK_CHECKPOINT_DIR, + IsValidSampleFn, ProduceBatchResult, ProduceBatchStatus, _TaskRunner, @@ -50,6 +51,7 @@ class DisaggTaskSpecConfig(BaseModel): weight: float = Field(default=1.0, ge=0.0) agent_loop_config: AgentLoopConfig judger_config: JudgerConfig | ComposedJudgerConfig | None = None + filter_func: IsValidSampleFn | None = None produce_strategy_config: DisaggProduceStrategyConfig = DisaggAsyncProduceStrategyConfig() sampler_config: SamplerConfig @@ -96,6 +98,7 @@ def build( agent_loop=agent_loop, produce_strategy=produce_strategy, sampler=sampler, + is_valid_sample_fn=task_cfg.filter_func, weight=task_cfg.weight, order=order, ) diff --git a/xtuner/v1/rl/agent_loop_manager/disagg_producer.py b/xtuner/v1/rl/agent_loop_manager/disagg_producer.py index d75c34b7c1..2bedd94e9f 100644 --- a/xtuner/v1/rl/agent_loop_manager/disagg_producer.py +++ b/xtuner/v1/rl/agent_loop_manager/disagg_producer.py @@ -13,14 +13,12 @@ from .produce_utils import ( PERIODIC_ABORT_INTERVAL_S, BaseProduceContext, - IsValidSampleFn, ProduceBatchStatus, ShouldContinueFn, _PendingTasks, _ProgressDisplayer, _put_claimed_tasks, calculate_stale_threshold, - default_is_valid_sample_fn, default_should_continue_fn, pause_pending_tasks, ) @@ -230,7 +228,6 @@ class DisaggProduceStrategyConfig(ABC, BaseModel): """非共卡后台 producer strategy 配置。""" model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) - is_valid_sample_fn: IsValidSampleFn = default_is_valid_sample_fn should_continue_fn: ShouldContinueFn = default_should_continue_fn @abstractmethod @@ -266,7 +263,6 @@ def build( max_staleness=self.max_staleness, sync_weights_interval=sync_weights_interval, tail_batch_trigger_size=self.tail_batch_trigger_size, - is_valid_sample_fn=self.is_valid_sample_fn, should_continue_fn=self.should_continue_fn, ) @@ -274,10 +270,8 @@ def build( class DisaggProduceStrategy(ABC): def __init__( self, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - self.is_valid_sample_fn = is_valid_sample_fn self.should_continue_fn = should_continue_fn @abstractmethod @@ -305,10 +299,9 @@ def __init__( tail_batch_trigger_size: int, max_staleness: int, sync_weights_interval: int, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - super().__init__(is_valid_sample_fn, should_continue_fn) + super().__init__(should_continue_fn) if not enable_partial_rollout and max_staleness > 0: logger.warning( @@ -382,7 +375,7 @@ async def produce_batch(self, ctx: DisaggProduceContext) -> ProduceBatchStatus: async def spawn_one() -> asyncio.Task: rollout_state = await ctx.sample_group(from_expired_pool=sample_from_expired) return create_task( - ctx.generate_group( + ctx.collect_rollout_group( rollout_state, enable_partial_rollout=self.enable_partial_rollout, ) diff --git a/xtuner/v1/rl/agent_loop_manager/produce_utils.py b/xtuner/v1/rl/agent_loop_manager/produce_utils.py index b80a2f9e99..a669fe2005 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -83,10 +83,6 @@ class ProduceBatchStatus(Enum): EXPIRED_BATCH = auto() -def default_is_valid_sample_fn(samples: list[RolloutState]) -> bool: - return True - - def default_should_continue_fn(completed_count: int, batch_size: int, **kwargs) -> bool: return completed_count < batch_size @@ -123,7 +119,7 @@ class BaseProduceContext: train_step: int model_step: int progress: "ProduceProgress | DisaggProduceProgress" - is_valid_sample_fn: IsValidSampleFn = default_is_valid_sample_fn + is_valid_sample_fn: IsValidSampleFn | None = None stale_threshold: int | None = None @property @@ -137,7 +133,7 @@ async def sample_group(self, *, from_expired_pool: bool) -> list[RolloutState]: group_status = [Status.EXPIRED, Status.ABORTED] if from_expired_pool else [Status.ABORTED] return await self.sampler.sample(task_name=self.task_name, group_status=group_status) - async def generate_group( + async def collect_rollout_group( self, rollout_state: list[RolloutState], *, @@ -151,13 +147,15 @@ async def generate_group( start = time.perf_counter() if isinstance(self.agent_loop, ray.actor.ActorHandle): - result = await self.agent_loop.generate_group.remote( + result = await self.agent_loop.collect_rollout_group.remote( rollout_state, + is_valid_sample_func=self.is_valid_sample_fn, enable_partial_rollout=enable_partial_rollout, ) else: - result = await self.agent_loop.generate_group( + result = await self.agent_loop.collect_rollout_group( rollout_state, + is_valid_sample_func=self.is_valid_sample_fn, enable_partial_rollout=enable_partial_rollout, ) elapsed = time.perf_counter() - start @@ -172,9 +170,8 @@ async def generate_group( async def put_generated_group(self, group: list[RolloutState]) -> bool: produced_tokens = sum(len(item.response_ids) for item in group if item.response_ids is not None) initial_status = get_group_status(group) - discard_status: Status | None = None - if initial_status == Status.COMPLETED: + if initial_status in (Status.COMPLETED, Status.FILTERED): rewards_sum = 0.0 rewards_count = 0 for item in group: @@ -188,15 +185,10 @@ async def put_generated_group(self, group: list[RolloutState]) -> bool: rewards_count += 1 self.progress.add_raw_rewards(self.task_name, rewards_sum, rewards_count) - if not self.is_valid_sample_fn(group): - discard_status = Status.FILTERED - elif initial_status == Status.FAILED: - discard_status = Status.FAILED - - if discard_status is not None: + if initial_status in (Status.FAILED, Status.FILTERED): # 失败样本和业务过滤样本都不进入 replay buffer。 self.progress.add_produced(self.task_name, samples=len(group), tokens=produced_tokens) - self.progress.add_discarded(self.task_name, discard_status, samples=len(group)) + self.progress.add_discarded(self.task_name, initial_status, samples=len(group)) for item in group: discard_rollout_state(item) return False @@ -272,13 +264,10 @@ class _TaskRunner: agent_loop: AgentLoopSpec produce_strategy: Any sampler: Sampler + is_valid_sample_fn: IsValidSampleFn | None = None weight: float = 1.0 order: int = 0 - @property - def is_valid_sample_fn(self) -> IsValidSampleFn: - return getattr(self.produce_strategy, "is_valid_sample_fn", default_is_valid_sample_fn) - @property def stale_threshold(self) -> int | None: return getattr(self.produce_strategy, "stale_threshold", None) diff --git a/xtuner/v1/rl/agent_loop_manager/producer.py b/xtuner/v1/rl/agent_loop_manager/producer.py index 5620be7772..f309e89ba5 100644 --- a/xtuner/v1/rl/agent_loop_manager/producer.py +++ b/xtuner/v1/rl/agent_loop_manager/producer.py @@ -13,12 +13,10 @@ from .produce_utils import ( PERIODIC_ABORT_INTERVAL_S, BaseProduceContext, - IsValidSampleFn, ShouldContinueFn, _ProgressDisplayer, _put_claimed_tasks, calculate_stale_threshold, - default_is_valid_sample_fn, default_should_continue_fn, pause_pending_tasks, ) @@ -127,16 +125,12 @@ class ProduceStrategyConfig(ABC, BaseModel): when it should stop producing samples for the current training step. Args: - is_valid_sample_fn (IsValidSampleFn): Function used to decide whether a - generated rollout group is trainable. Defaults to - ``default_is_valid_sample_fn``. should_continue_fn (ShouldContinueFn): Function used to decide whether production should continue after a group is processed. Defaults to ``default_should_continue_fn``. """ model_config = ConfigDict(extra="forbid", arbitrary_types_allowed=True) - is_valid_sample_fn: IsValidSampleFn = default_is_valid_sample_fn should_continue_fn: ShouldContinueFn = default_should_continue_fn @abstractmethod @@ -156,9 +150,6 @@ class SyncProduceStrategyConfig(ProduceStrategyConfig): in a colocated or tightly synchronized workflow. Args: - is_valid_sample_fn (IsValidSampleFn): Function used to decide whether a - generated rollout group is trainable. Defaults to - ``default_is_valid_sample_fn``. should_continue_fn (ShouldContinueFn): Function used to decide whether production should continue after a group is processed. Defaults to ``default_should_continue_fn``. @@ -176,9 +167,7 @@ def build( sync_weights_interval: int = 1, rollout_controller: "Optional[RolloutControllerProxy]" = None, ) -> "SyncProduceStrategy": - return SyncProduceStrategy( - is_valid_sample_fn=self.is_valid_sample_fn, should_continue_fn=self.should_continue_fn - ) + return SyncProduceStrategy(should_continue_fn=self.should_continue_fn) class AsyncProduceStrategyConfig(ProduceStrategyConfig): @@ -190,9 +179,6 @@ class AsyncProduceStrategyConfig(ProduceStrategyConfig): discard samples that are too stale relative to the current training step. Args: - is_valid_sample_fn (IsValidSampleFn): Function used to decide whether a - generated rollout group is trainable. Defaults to - ``default_is_valid_sample_fn``. should_continue_fn (ShouldContinueFn): Function used to decide whether production should continue after a group is processed. Defaults to ``default_should_continue_fn``. @@ -237,7 +223,6 @@ def build( max_staleness=self.max_staleness, sync_weights_interval=sync_weights_interval, tail_batch_trigger_size=self.tail_batch_trigger_size, - is_valid_sample_fn=self.is_valid_sample_fn, should_continue_fn=self.should_continue_fn, ) @@ -245,10 +230,8 @@ def build( class ProduceStrategy(ABC): def __init__( self, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - self.is_valid_sample_fn = is_valid_sample_fn self.should_continue_fn = should_continue_fn @abstractmethod @@ -268,7 +251,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: for _ in range(ctx.task_batch_size): rollout_state = await ctx.sampler.sample(task_name=ctx.task_name) - task = create_task(ctx.generate_group(rollout_state)) + task = create_task(ctx.collect_rollout_group(rollout_state)) pending_tasks.add(task) logger.info(f"[SyncProduceStrategy] Started {len(pending_tasks)} initial tasks.") @@ -286,7 +269,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: done_tasks, pending_tasks = await asyncio.wait( pending_tasks, timeout=1, return_when=asyncio.FIRST_COMPLETED ) - # put_generated_group 负责过滤和入库。 + # AgentLoop 已完成过滤;put_generated_group 只处理状态、数据入库和释放。 for task in done_tasks: items = task.result() @@ -301,7 +284,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: completed_sample_count, ctx.task_batch_size ): rollout_state = await ctx.sampler.sample(task_name=ctx.task_name) - task = create_task(ctx.generate_group(rollout_state)) + task = create_task(ctx.collect_rollout_group(rollout_state)) pending_tasks.add(task) progress_displayer.close() @@ -316,10 +299,9 @@ def __init__( tail_batch_trigger_size: int, max_staleness: int, sync_weights_interval: int, - is_valid_sample_fn: IsValidSampleFn, should_continue_fn: ShouldContinueFn, ): - super().__init__(is_valid_sample_fn, should_continue_fn) + super().__init__(should_continue_fn) # TODO: 需要添加 tail_batch_max_tries # 作用是:如果一个样本多次重试,则将它置为特殊状态 MAX_TRIES,这类样本和过期样本一起触发tail batch逻辑 @@ -381,7 +363,7 @@ async def produce_batch(self, ctx: ProduceContext) -> None: async def spawn_one() -> asyncio.Task: rollout_state = await ctx.sample_group(from_expired_pool=sample_from_expired) return create_task( - ctx.generate_group( + ctx.collect_rollout_group( rollout_state, enable_partial_rollout=self.enable_partial_rollout, ) From aeaa1d079eb4bc77441769b8a21dcb2cbe97870f Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Wed, 22 Jul 2026 05:45:47 +0000 Subject: [PATCH 2/8] [Feat] Add Teacher logprob client to AgentLoop --- tests/rl/test_on_policy_distillation.py | 449 ++++++++++++++++++++++++ tests/rl/test_producer.py | 1 + xtuner/v1/data_proto/rl_data.py | 4 + xtuner/v1/rl/agent_loop/agent_loop.py | 75 +++- xtuner/v1/rl/on_policy_distillation.py | 129 +++++++ 5 files changed, 656 insertions(+), 2 deletions(-) create mode 100644 tests/rl/test_on_policy_distillation.py create mode 100644 xtuner/v1/rl/on_policy_distillation.py diff --git a/tests/rl/test_on_policy_distillation.py b/tests/rl/test_on_policy_distillation.py new file mode 100644 index 0000000000..1485416d7a --- /dev/null +++ b/tests/rl/test_on_policy_distillation.py @@ -0,0 +1,449 @@ +import asyncio +import json +import socket +import threading +import time +import unittest +from collections.abc import Callable +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any +from unittest.mock import MagicMock + +from xtuner.v1.data_proto.rl_data import ( + RolloutState, + SampleParams, + Status, + get_group_status, + reset_rollout_response, +) +from xtuner.v1.rl.agent_loop.single_turn_agent_loop import SingleTurnAgentLoop +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + OPDTeacherConfig, + TeacherLogprobClient, +) +from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig + + +@dataclass +class _HTTPResponse: + body: Any + status: int = 200 + delay_s: float = 0.0 + + +class _EndpointState: + def __init__( + self, + responses: list[_HTTPResponse] | None = None, + on_request: Callable[[dict[str, Any]], None] | None = None, + ) -> None: + self.responses = list(responses or []) + self.on_request = on_request + self.requests: list[dict[str, Any]] = [] + self.headers: list[dict[str, str]] = [] + self.request_event = threading.Event() + self.max_active_requests = 0 + self._active_requests = 0 + self._lock = threading.Lock() + + def handle(self, payload: dict[str, Any], headers: dict[str, str]) -> _HTTPResponse: + with self._lock: + self.requests.append(payload) + self.headers.append(headers) + self._active_requests += 1 + self.max_active_requests = max(self.max_active_requests, self._active_requests) + response = self.responses.pop(0) if self.responses else self._success_response(payload) + self.request_event.set() + if self.on_request is not None: + self.on_request(payload) + try: + if response.delay_s: + time.sleep(response.delay_s) + return response + finally: + with self._lock: + self._active_requests -= 1 + + @staticmethod + def _success_response(payload: dict[str, Any]) -> _HTTPResponse: + scored_tokens = payload["input_ids"][payload["logprob_start_len"] :] + return _HTTPResponse( + body={ + "meta_info": { + "input_token_logprobs": [[-0.1 - index / 100, token] for index, token in enumerate(scored_tokens)] + } + } + ) + + +class _Handler(BaseHTTPRequestHandler): + def do_POST(self) -> None: + content_length = int(self.headers["Content-Length"]) + payload = json.loads(self.rfile.read(content_length)) + response = self.server.endpoint_state.handle(payload, dict(self.headers)) # type: ignore[attr-defined] + body = response.body if isinstance(response.body, bytes) else json.dumps(response.body).encode() + self.send_response(response.status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + try: + self.wfile.write(body) + except (BrokenPipeError, ConnectionResetError): + pass + + def log_message(self, format: str, *args: Any) -> None: + return + + +class _FakeEndpoint: + def __init__( + self, + responses: list[_HTTPResponse] | None = None, + on_request: Callable[[dict[str, Any]], None] | None = None, + ) -> None: + self.state = _EndpointState(responses, on_request) + self.server = ThreadingHTTPServer(("127.0.0.1", 0), _Handler) + self.server.endpoint_state = self.state # type: ignore[attr-defined] + self.thread = threading.Thread(target=self.server.serve_forever, daemon=True) + self.thread.start() + + @property + def url(self) -> str: + host, port = self.server.server_address + return f"http://{host}:{port}" + + def close(self) -> None: + self.server.shutdown() + self.server.server_close() + self.thread.join(timeout=2) + + +def _state( + rollout_id: int, + *, + data_source: str = "math", + prompt_ids: list[int] | None = None, + response_ids: list[int] | None = None, +) -> RolloutState: + prompt_ids = prompt_ids or [100 + rollout_id] + response_ids = response_ids or [200 + rollout_id, 2] + return RolloutState( + rollout_id=rollout_id, + group_id=0, + message=[{"role": "user", "content": f"prompt {rollout_id}"}], + prompt_ids=prompt_ids, + tokens=prompt_ids, + response="response", + response_ids=response_ids, + logprobs=[-0.4] * len(response_ids), + response_mask=[1] * len(response_ids), + status=Status.COMPLETED, + extra_fields={"origin_data_source": data_source}, + ) + + +def _opd_config(endpoint_by_name: dict[str, str], data_source_teacher_map: dict[str, str], **kwargs) -> OPDConfig: + return OPDConfig( + teachers=[ + OPDTeacherConfig(name=name, endpoint=endpoint, **kwargs) for name, endpoint in endpoint_by_name.items() + ], + data_source_teacher_map=data_source_teacher_map, + ) + + +class TestTeacherLogprobClient(unittest.IsolatedAsyncioTestCase): + def _endpoint(self, responses: list[_HTTPResponse] | None = None) -> _FakeEndpoint: + endpoint = _FakeEndpoint(responses) + self.addCleanup(endpoint.close) + return endpoint + + async def test_compute_logprobs_uses_sglang_prefill_protocol(self): + endpoint = self._endpoint() + client = TeacherLogprobClient( + OPDTeacherConfig(name="teacher", endpoint=endpoint.url, api_key="secret", max_retry_per_sample=0) + ) + self.addAsyncCleanup(client.close) + state = _state(1, prompt_ids=[10], response_ids=[20, 2]) + + result = await client.compute_logprobs(state) + + self.assertIs(result, state) + self.assertEqual(result.teacher_tokens, [20, 2]) + for actual, expected in zip(result.teacher_logprobs or [], [-0.11, -0.12], strict=True): + self.assertAlmostEqual(actual, expected) + payload = endpoint.state.requests[0] + self.assertEqual(payload["input_ids"], [10, 20, 2]) + self.assertEqual(payload["logprob_start_len"], 0) + self.assertEqual(payload["top_logprobs_num"], 0) + self.assertEqual( + payload["sampling_params"], + {"max_new_tokens": 0, "temperature": 1.0, "skip_special_tokens": False}, + ) + self.assertEqual(endpoint.state.headers[0]["Authorization"], "Bearer secret") + + async def test_invalid_teacher_responses_fail_without_signal(self): + responses = [ + _HTTPResponse(b"not-json"), + _HTTPResponse({"missing": "meta_info"}), + _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [-0.2, 20]]}}), + _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [-0.2, 20], [-0.3, 2], [-0.4, 3]]}}), + _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [-0.2, 999], [-0.3, 2]]}}), + _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [float("nan"), 20], [-0.3, 2]]}}), + _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [float("inf"), 20], [-0.3, 2]]}}), + ] + endpoint = self._endpoint(responses) + client = TeacherLogprobClient(OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=0)) + self.addAsyncCleanup(client.close) + + for rollout_id in range(len(responses)): + with self.subTest(rollout_id=rollout_id): + result = await client.compute_logprobs(_state(rollout_id, prompt_ids=[10], response_ids=[20, 2])) + self.assertEqual(result.status, Status.FAILED) + self.assertIsNone(result.teacher_tokens) + self.assertIsNone(result.teacher_logprobs) + self.assertIn("scoring failed after 1 attempts", result.error_msg or "") + + async def test_http_status_retries_then_succeeds(self): + endpoint = self._endpoint([_HTTPResponse({"error": "busy"}, status=503)]) + client = TeacherLogprobClient(OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=1)) + self.addAsyncCleanup(client.close) + + result = await client.compute_logprobs(_state(1)) + + self.assertEqual(result.status, Status.COMPLETED) + self.assertEqual(len(endpoint.state.requests), 2) + + async def test_http_status_retry_exhaustion_marks_failed(self): + for status in (400, 500): + with self.subTest(status=status): + endpoint = self._endpoint([_HTTPResponse({"error": "failed"}, status=status) for _ in range(2)]) + client = TeacherLogprobClient( + OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=1) + ) + self.addAsyncCleanup(client.close) + + result = await client.compute_logprobs(_state(status)) + + self.assertEqual(result.status, Status.FAILED) + self.assertEqual(len(endpoint.state.requests), 2) + + async def test_timeout_and_connection_error_use_bounded_retries(self): + endpoint = self._endpoint([_HTTPResponse({}, delay_s=0.05) for _ in range(2)]) + timeout_client = TeacherLogprobClient( + OPDTeacherConfig( + name="timeout", + endpoint=endpoint.url, + request_timeout_s=0.01, + max_retry_per_sample=1, + ) + ) + self.addAsyncCleanup(timeout_client.close) + + timeout_result = await timeout_client.compute_logprobs(_state(1)) + + self.assertEqual(timeout_result.status, Status.FAILED) + self.assertEqual(len(endpoint.state.requests), 2) + + sock = socket.socket() + sock.bind(("127.0.0.1", 0)) + closed_port = sock.getsockname()[1] + sock.close() + connection_client = TeacherLogprobClient( + OPDTeacherConfig( + name="connection", + endpoint=f"http://127.0.0.1:{closed_port}", + request_timeout_s=0.1, + max_retry_per_sample=1, + ) + ) + self.addAsyncCleanup(connection_client.close) + + connection_result = await connection_client.compute_logprobs(_state(2)) + + self.assertEqual(connection_result.status, Status.FAILED) + + async def test_max_concurrency_limits_requests_per_client(self): + endpoint = self._endpoint([_HTTPResponse({}, delay_s=0.05) for _ in range(5)]) + client = TeacherLogprobClient( + OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=0, max_concurrency=2) + ) + self.addAsyncCleanup(client.close) + + results = await asyncio.gather(*(client.compute_logprobs(_state(index)) for index in range(5))) + + self.assertEqual(endpoint.state.max_active_requests, 2) + self.assertTrue(all(result.status == Status.FAILED for result in results)) + + +class TestAgentLoopTeacherScoring(unittest.IsolatedAsyncioTestCase): + def _endpoint( + self, + responses: list[_HTTPResponse] | None = None, + on_request: Callable[[dict[str, Any]], None] | None = None, + ) -> _FakeEndpoint: + endpoint = _FakeEndpoint(responses, on_request) + self.addCleanup(endpoint.close) + return endpoint + + def _loop(self, opd_config: OPDConfig) -> SingleTurnAgentLoop: + loop = SingleTurnAgentLoop.__new__(SingleTurnAgentLoop) + loop.rollout_ctl = MagicMock() + loop.sample_params = SampleParams(max_tokens=8) + loop.judger = None + loop.enable_batch_judge = False + loop._judger_pause_event = asyncio.Event() + loop.logger = MagicMock() + loop.configure_opd(opd_config) + self.addAsyncCleanup(loop.close) + return loop + + @staticmethod + def _complete(state: RolloutState) -> RolloutState: + state.status = Status.COMPLETED + return state + + async def test_eager_scoring_starts_before_other_samples_finish_generation(self): + endpoint = self._endpoint() + loop = self._loop(_opd_config({"teacher": endpoint.url}, {"math": "teacher"})) + release_second = asyncio.Event() + second_state = _state(2) + second_state.status = Status.INIT + + async def generate_sample(state, **kwargs): + if state is second_state: + await release_second.wait() + return self._complete(state) + + loop.generate_sample = generate_sample + task = asyncio.create_task(loop.collect_rollout_group([_state(1), second_state])) + + request_started = await asyncio.to_thread(endpoint.state.request_event.wait, 1.0) + self.assertTrue(request_started) + self.assertEqual(second_state.status, Status.INIT) + release_second.set() + result = await task + + self.assertEqual(get_group_status(result), Status.COMPLETED) + self.assertEqual(len(endpoint.state.requests), 2) + self.assertTrue(all(state.teacher_logprobs is not None for state in result)) + self.assertTrue(all("teacher_score_time_s" in state.extra_fields for state in result)) + + async def test_filter_false_skips_teacher_endpoint(self): + endpoint = self._endpoint() + loop = self._loop(_opd_config({"teacher": endpoint.url}, {"math": "teacher"})) + + async def generate_sample(state, **kwargs): + return self._complete(state) + + loop.generate_sample = generate_sample + result = await loop.collect_rollout_group( + [_state(1), _state(2)], + is_valid_sample_func=lambda group: False, + ) + + self.assertEqual(get_group_status(result), Status.FILTERED) + self.assertEqual(endpoint.state.requests, []) + self.assertTrue(all("teacher_score_time_s" not in state.extra_fields for state in result)) + + async def test_filter_true_runs_before_lazy_teacher_scoring(self): + order = [] + endpoint = self._endpoint(on_request=lambda payload: order.append("teacher")) + + def filter_func(group): + order.append("filter") + return True + + loop = self._loop(_opd_config({"teacher": endpoint.url}, {"math": "teacher"})) + + async def generate_sample(state, **kwargs): + return self._complete(state) + + loop.generate_sample = generate_sample + result = await loop.collect_rollout_group([_state(1)], is_valid_sample_func=filter_func) + + self.assertEqual(order, ["filter", "teacher"]) + self.assertEqual(result[0].status, Status.COMPLETED) + self.assertIsNotNone(result[0].teacher_logprobs) + self.assertIn("teacher_score_time_s", result[0].extra_fields) + + async def test_origin_data_source_routes_groups_to_different_teachers(self): + math_endpoint = self._endpoint() + code_endpoint = self._endpoint() + loop = self._loop( + _opd_config( + {"math_teacher": math_endpoint.url, "code_teacher": code_endpoint.url}, + {"math": "math_teacher", "code": "code_teacher"}, + ) + ) + + async def generate_sample(state, **kwargs): + return self._complete(state) + + loop.generate_sample = generate_sample + await loop.collect_rollout_group([_state(1, data_source="math")]) + await loop.collect_rollout_group([_state(2, data_source="code")]) + + self.assertEqual(len(math_endpoint.state.requests), 1) + self.assertEqual(len(code_endpoint.state.requests), 1) + + async def test_teacher_failure_marks_group_failed(self): + endpoint = self._endpoint([_HTTPResponse({"error": "down"}, status=500)]) + loop = self._loop( + _opd_config( + {"teacher": endpoint.url}, + {"math": "teacher"}, + max_retry_per_sample=0, + ) + ) + + async def generate_sample(state, **kwargs): + return self._complete(state) + + loop.generate_sample = generate_sample + result = await loop.collect_rollout_group([_state(1)]) + + self.assertEqual(get_group_status(result), Status.FAILED) + self.assertIsNone(result[0].teacher_logprobs) + self.assertIn("teacher_score_time_s", result[0].extra_fields) + + +class TestOPDRolloutData(unittest.IsolatedAsyncioTestCase): + async def test_rollout_state_round_trip_replay_and_reset_preserve_contract(self): + state = _state(1) + state.teacher_tokens = [201, 2] + state.teacher_logprobs = [-0.2, -0.3] + + restored = RolloutState.model_validate(state.model_dump()) + self.assertEqual(restored.teacher_tokens, state.teacher_tokens) + self.assertEqual(restored.teacher_logprobs, state.teacher_logprobs) + + replay_buffer = AsyncReplayBufferConfig().build() + await replay_buffer.put([restored], "math") + replayed = (await replay_buffer.get(1, "math", Status.COMPLETED))[0][0] + self.assertEqual(replayed.teacher_tokens, [201, 2]) + self.assertEqual(replayed.teacher_logprobs, [-0.2, -0.3]) + + reset_rollout_response(replayed) + self.assertIsNone(replayed.teacher_tokens) + self.assertIsNone(replayed.teacher_logprobs) + + +class TestOPDConfig(unittest.TestCase): + def test_rejects_duplicate_or_unknown_teacher_names(self): + teacher = OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:1") + with self.assertRaisesRegex(ValueError, "must be unique"): + OPDConfig( + teachers=[teacher, teacher], + data_source_teacher_map={"math": "teacher"}, + ) + with self.assertRaisesRegex(ValueError, "unknown teachers"): + OPDConfig( + teachers=[teacher], + data_source_teacher_map={"math": "missing"}, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index b20815d81a..c1c187779b 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -118,6 +118,7 @@ async def mock_gen(rs, **kwargs): return rs mock_agent_loop.generate_group = mock_gen + mock_agent_loop.teacher_clients = {} mock_agent_loop.collect_rollout_group = AgentLoop.collect_rollout_group.__get__(mock_agent_loop) return mock_agent_loop diff --git a/xtuner/v1/data_proto/rl_data.py b/xtuner/v1/data_proto/rl_data.py index 7db20f2940..fb22c635be 100644 --- a/xtuner/v1/data_proto/rl_data.py +++ b/xtuner/v1/data_proto/rl_data.py @@ -109,6 +109,8 @@ class RolloutState(BaseModel): tool_calls: list[RolloutToolCall] | None = None response_ids: list[int] | None = None logprobs: list[float] | None = None + teacher_tokens: list[int] | None = None + teacher_logprobs: list[float] | 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 中覆盖写 @@ -244,6 +246,8 @@ def reset_rollout_response(rollout_state: RolloutState) -> RolloutState: rollout_state.response = "" rollout_state.response_ids = [] rollout_state.logprobs = [] + rollout_state.teacher_tokens = None + rollout_state.teacher_logprobs = None rollout_state.routed_experts = None rollout_state.finish_reason = None rollout_state.response_mask = [] diff --git a/xtuner/v1/rl/agent_loop/agent_loop.py b/xtuner/v1/rl/agent_loop/agent_loop.py index 22c23186f9..754087af96 100644 --- a/xtuner/v1/rl/agent_loop/agent_loop.py +++ b/xtuner/v1/rl/agent_loop/agent_loop.py @@ -12,6 +12,11 @@ from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status, get_group_status from xtuner.v1.rl.judger import Judger +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + TeacherLogprobClient, + route_teacher_client, +) from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.rl.rollout.constants import AGENT_LOOP_RAY_GENERATE_MAX_CONCURRENCY from xtuner.v1.rl.trace.rollout_api import ( @@ -40,13 +45,22 @@ class AgentLoopConfig(ABC, BaseModel): enable_batch_judge: bool = False requires_rollout_proxy: bool = False - def build(self, rollout_controller, judger: Judger | None = None, logger=None) -> AgentLoopSpec: + def build( + self, + rollout_controller, + judger: Judger | None = None, + logger=None, + *, + opd_config: OPDConfig | None = None, + ) -> AgentLoopSpec: if self.cpu_resources is None: - return self.build_local( + agent_loop = self.build_local( rollout_controller=rollout_controller, judger=judger, logger=logger, ) + agent_loop.configure_opd(opd_config) + return agent_loop concurrency = AGENT_LOOP_RAY_GENERATE_MAX_CONCURRENCY @@ -62,6 +76,7 @@ def build(self, rollout_controller, judger: Judger | None = None, logger=None) - concurrency=concurrency, judger=judger, logger=logger, + opd_config=opd_config, ) return self._build_ray_actor( rollout_controller=rollout_controller, @@ -69,6 +84,7 @@ def build(self, rollout_controller, judger: Judger | None = None, logger=None) - concurrency=concurrency, judger=judger, logger=logger, + opd_config=opd_config, ) @abstractmethod @@ -87,6 +103,7 @@ def _build_ray_actor( pg: PlacementGroup | None = None, judger: Judger | None = None, logger=None, + opd_config: OPDConfig | None = None, ) -> RayAgentLoopProxy: ray_agent_loop = ray.remote( concurrency_groups={ @@ -105,6 +122,7 @@ def _build_ray_actor( actor_num_cpus=cpu_resources.num_cpus_per_worker, actor_memory=cpu_resources.cpu_memory_per_worker, capture_child_tasks=True, + opd_config=opd_config, ), ) @@ -117,6 +135,7 @@ def _build_ray_actors( judger: Judger | None = None, logger=None, start_bundle_idx: int = 0, + opd_config: OPDConfig | None = None, ) -> list[RayAgentLoopProxy]: ray_agent_loop = ray.remote( concurrency_groups={ @@ -136,6 +155,7 @@ def _build_ray_actors( actor_num_cpus_per_worker=cpu_resources.num_cpus_per_worker, actor_memory_per_worker=cpu_resources.cpu_memory_per_worker, capture_child_tasks=True, + opd_config=opd_config, ), ) @@ -148,6 +168,7 @@ def _build_router( judger: Judger | None = None, logger=None, start_bundle_idx: int = 0, + opd_config: OPDConfig | None = None, ) -> RouterAgentLoop: return RouterAgentLoop( workers=self._build_ray_actors( @@ -158,6 +179,7 @@ def _build_router( judger=judger, logger=logger, start_bundle_idx=start_bundle_idx, + opd_config=opd_config, ), rollout_ctl=rollout_controller, ) @@ -185,6 +207,15 @@ def __init__( else: self.logger = logger self._judger_pause_event = asyncio.Event() + self.teacher_clients: dict[str, TeacherLogprobClient] = {} + self.data_source_teacher_map: dict[str, str] = {} + + def configure_opd(self, opd_config: OPDConfig | None) -> None: + if opd_config is None: + return + + self.teacher_clients = {teacher.name: TeacherLogprobClient(teacher) for teacher in opd_config.teachers} + self.data_source_teacher_map = dict(opd_config.data_source_teacher_map) @abstractmethod async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: ... @@ -209,12 +240,39 @@ async def collect_rollout_group( is_valid_sample_func: Callable[[list[RolloutState]], bool] | None = None, **kwargs, ) -> list[RolloutState]: + if is_valid_sample_func is None and self.teacher_clients: + teacher = route_teacher_client( + rollout_state[0], + data_source_teacher_map=self.data_source_teacher_map, + teacher_clients=self.teacher_clients, + ) + + async def generate_and_score(state: RolloutState) -> RolloutState: + state.sample_params = self.sample_params + state = await self.generate_sample(state, **kwargs) + if state.status == Status.COMPLETED: + state = await teacher.compute_logprobs(state) + return state + + group = list(await asyncio.gather(*(create_task(generate_and_score(state)) for state in rollout_state))) + if self.judger is not None and self.enable_batch_judge and get_group_status(group) == Status.COMPLETED: + group = await self.run_judger(group) + return group + group = await self.generate_group(rollout_state, **kwargs) if get_group_status(group) != Status.COMPLETED: return group if is_valid_sample_func is not None and not is_valid_sample_func(group): for state in group: state.status = Status.FILTERED + return group + if self.teacher_clients: + teacher = route_teacher_client( + group[0], + data_source_teacher_map=self.data_source_teacher_map, + teacher_clients=self.teacher_clients, + ) + group = list(await asyncio.gather(*(create_task(teacher.compute_logprobs(state)) for state in group))) return group @overload @@ -267,6 +325,9 @@ async def pause(self) -> None: finally: self._judger_pause_event.clear() + async def close(self) -> None: + await asyncio.gather(*(client.close() for client in self.teacher_clients.values())) + class RouterAgentLoop: def __init__(self, workers: list[RayAgentLoopProxy], rollout_ctl: RolloutController): @@ -318,6 +379,9 @@ async def pause(self) -> None: *(worker.pause.remote() for worker in self.workers), ) + async def close(self) -> None: + await asyncio.gather(*(worker.close.remote() for worker in self.workers)) + async def get_agent_loop_rollout_ctl(agent_loop: AgentLoopSpec) -> RolloutController: rollout_ctl = getattr(agent_loop, "rollout_ctl", None) @@ -337,12 +401,15 @@ def __init__( rollout_controller: RolloutController, judger: Judger | None = None, logger=None, + *, + opd_config: OPDConfig | None = None, ): self.agent_loop = agent_loop_config.build_local( rollout_controller=rollout_controller, judger=judger, logger=logger, ) + self.agent_loop.configure_opd(opd_config) @ray_method(concurrency_group=AGENT_LOOP_CONCURRENCY_GROUP_GENERATE) async def generate_sample(self, rollout_state: RolloutState, **kwargs) -> RolloutState: @@ -364,6 +431,10 @@ async def get_rollout_ctl(self): async def pause(self) -> None: return await self.agent_loop.pause() + @ray_method + async def close(self) -> None: + return await self.agent_loop.close() + RayAgentLoop = cast( ActorClass[AgentLoopActor], diff --git a/xtuner/v1/rl/on_policy_distillation.py b/xtuner/v1/rl/on_policy_distillation.py new file mode 100644 index 0000000000..69f6118dfa --- /dev/null +++ b/xtuner/v1/rl/on_policy_distillation.py @@ -0,0 +1,129 @@ +from __future__ import annotations + +import asyncio +import math +import time +from typing import Literal, cast + +import httpx +from pydantic import BaseModel, ConfigDict, Field, model_validator + +from xtuner.v1.data_proto.rl_data import RolloutState, Status + + +class OPDTeacherConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + name: str + endpoint: str + api_key: str | None = None + request_timeout_s: float = Field(default=1200.0, gt=0.0) + max_retry_per_sample: int = Field(default=2, ge=0) + max_concurrency: int = Field(default=128, gt=0) + + +class OPDConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + mode: Literal["pg-opd"] = "pg-opd" + task_adv_weight: float = Field(default=0.0, ge=0.0) + opd_adv_weight: float = Field(default=1.0, ge=0.0) + teachers: list[OPDTeacherConfig] = Field(min_length=1) + data_source_teacher_map: dict[str, str] = Field(min_length=1) + + @model_validator(mode="after") + def validate_teacher_names(self) -> OPDConfig: + teacher_names = [teacher.name for teacher in self.teachers] + if len(teacher_names) != len(set(teacher_names)): + raise ValueError("OPD teacher names must be unique") + unknown_teachers = set(self.data_source_teacher_map.values()) - set(teacher_names) + if unknown_teachers: + raise ValueError(f"data_source_teacher_map references unknown teachers: {sorted(unknown_teachers)}") + return self + + +class TeacherLogprobClient: + """Minimal asynchronous client for one external SGLang teacher.""" + + def __init__(self, config: OPDTeacherConfig) -> None: + self.config = config + self.name = config.name + self.url = f"{config.endpoint.rstrip('/')}/generate" + self._semaphore = asyncio.Semaphore(config.max_concurrency) + + headers = {"Content-Type": "application/json"} + if config.api_key is not None: + headers["Authorization"] = f"Bearer {config.api_key}" + self._client = httpx.AsyncClient(headers=headers, timeout=config.request_timeout_s) + + async def compute_logprobs(self, state: RolloutState) -> RolloutState: + start = time.perf_counter() + try: + prompt_ids = cast(list[int], state.prompt_ids) + response_ids = cast(list[int], state.response_ids) + payload = { + "input_ids": prompt_ids + response_ids, + "sampling_params": { + "max_new_tokens": 0, + "temperature": 1.0, + "skip_special_tokens": False, + }, + "return_logprob": True, + "logprob_start_len": max(len(prompt_ids) - 1, 0), + "top_logprobs_num": 0, + "stream": False, + } + + retries = 0 + while True: + try: + async with self._semaphore: + response = await self._client.post(self.url, json=payload) + response.raise_for_status() + teacher_tokens, teacher_logprobs = self._parse_response(response, response_ids) + state.teacher_tokens = teacher_tokens + state.teacher_logprobs = teacher_logprobs + return state + except (httpx.HTTPStatusError, httpx.RequestError, ValueError) as exc: + if retries >= self.config.max_retry_per_sample: + state.status = Status.FAILED + state.error_msg = f"Teacher {self.name!r} scoring failed after {retries + 1} attempts: {exc}" + return state + retries += 1 + await asyncio.sleep(0.1) + finally: + state.extra_fields["teacher_score_time_s"] = time.perf_counter() - start + + def _parse_response( + self, + response: httpx.Response, + response_ids: list[int], + ) -> tuple[list[int], list[float]]: + try: + raw_logprobs = response.json()["meta_info"]["input_token_logprobs"] + teacher_tokens = [item[1] for item in raw_logprobs[1:]] + teacher_logprobs = [float(item[0]) for item in raw_logprobs[1:]] + except (KeyError, TypeError, IndexError) as exc: + raise ValueError("Invalid teacher response") from exc + + if len(raw_logprobs) != len(response_ids) + 1: + raise ValueError("Teacher logprob length mismatch") + if teacher_tokens != response_ids: + raise ValueError("Teacher token ids mismatch") + if not all(math.isfinite(logprob) for logprob in teacher_logprobs): + raise ValueError("Teacher logprobs contain NaN or Inf") + return teacher_tokens, teacher_logprobs + + async def close(self) -> None: + await self._client.aclose() + + +def route_teacher_client( + state: RolloutState, + *, + data_source_teacher_map: dict[str, str], + teacher_clients: dict[str, TeacherLogprobClient], +) -> TeacherLogprobClient: + data_source = state.extra_fields["origin_data_source"] + teacher_name = data_source_teacher_map[data_source] + return teacher_clients[teacher_name] From 7df6d4f12121dbf9ab5a8c770c021f4da7c0636a Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Wed, 22 Jul 2026 06:36:10 +0000 Subject: [PATCH 3/8] [Feat] Add PG-OPD token advantages --- tests/rl/test_on_policy_distillation.py | 266 ++++++++++++++++++++++++ tests/rl/test_prepare_train_data.py | 75 ++++++- xtuner/v1/rl/on_policy_distillation.py | 53 +++++ xtuner/v1/train/rl_trainer.py | 104 +++++---- 4 files changed, 459 insertions(+), 39 deletions(-) diff --git a/tests/rl/test_on_policy_distillation.py b/tests/rl/test_on_policy_distillation.py index 1485416d7a..937c169be8 100644 --- a/tests/rl/test_on_policy_distillation.py +++ b/tests/rl/test_on_policy_distillation.py @@ -1,5 +1,6 @@ import asyncio import json +import math import socket import threading import time @@ -10,6 +11,8 @@ from typing import Any from unittest.mock import MagicMock +import torch + from xtuner.v1.data_proto.rl_data import ( RolloutState, SampleParams, @@ -17,13 +20,18 @@ get_group_status, reset_rollout_response, ) +from xtuner.v1.data_proto.sequence_context import SequenceContext from xtuner.v1.rl.agent_loop.single_turn_agent_loop import SingleTurnAgentLoop +from xtuner.v1.rl.loss import GRPOLossConfig, GRPOLossContext from xtuner.v1.rl.on_policy_distillation import ( OPDConfig, OPDTeacherConfig, TeacherLogprobClient, + compute_pg_opd_token_advantages, ) from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig +from xtuner.v1.rl.rollout_is import RolloutImportanceSampling +from xtuner.v1.rl.trainer.controller import TrainingController @dataclass @@ -153,6 +161,25 @@ def _opd_config(endpoint_by_name: dict[str, str], data_source_teacher_map: dict[ ) +def _algorithm_config(*, task_adv_weight: float = 0.0, opd_adv_weight: float = 1.0) -> OPDConfig: + return OPDConfig( + task_adv_weight=task_adv_weight, + opd_adv_weight=opd_adv_weight, + teachers=[OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:1")], + data_source_teacher_map={"math": "teacher"}, + ) + + +class _FixedAdvantageEstimator: + def __init__(self, values: list[float]) -> None: + self.values = values + self.calls: list[tuple[torch.Tensor, list[RolloutState]]] = [] + + def compute(self, rewards: torch.Tensor, group: list[RolloutState]) -> torch.Tensor: + self.calls.append((rewards.clone(), group)) + return torch.tensor(self.values[: len(group)], dtype=torch.float32) + + class TestTeacherLogprobClient(unittest.IsolatedAsyncioTestCase): def _endpoint(self, responses: list[_HTTPResponse] | None = None) -> _FakeEndpoint: endpoint = _FakeEndpoint(responses) @@ -430,6 +457,245 @@ async def test_rollout_state_round_trip_replay_and_reset_preserve_contract(self) self.assertIsNone(replayed.teacher_logprobs) +class TestPGOPDTokenAdvantages(unittest.TestCase): + @staticmethod + def _scored_state( + rollout_id: int, + *, + behavior_logprobs: list[float], + teacher_logprobs: list[float], + reward: float | None = None, + response_mask: list[int] | None = None, + ) -> RolloutState: + response_ids = list(range(100, 100 + len(behavior_logprobs))) + state = _state(rollout_id, response_ids=response_ids) + state.logprobs = behavior_logprobs + state.teacher_tokens = response_ids + state.teacher_logprobs = teacher_logprobs + state.reward = None if reward is None else {"score": reward} + state.response_mask = response_mask + return state + + def test_pure_opd_uses_token_delta_and_response_mask_without_reward(self): + state = self._scored_state( + 1, + behavior_logprobs=[-1.0, -2.0, -3.0], + teacher_logprobs=[-0.5, -2.5, -3.0], + response_mask=[1, 0, 1], + ) + + advantages = compute_pg_opd_token_advantages( + [state], + config=_algorithm_config(task_adv_weight=0.0, opd_adv_weight=2.0), + task_adv_estimator=None, + ) + + torch.testing.assert_close(advantages[0], torch.tensor([1.0, 0.0, 0.0])) + self.assertFalse(advantages[0].requires_grad) + + state.response_mask = [] + unmasked = compute_pg_opd_token_advantages( + [state], + config=_algorithm_config(), + task_adv_estimator=None, + ) + torch.testing.assert_close(unmasked[0], torch.tensor([0.5, -0.5, 0.0])) + + def test_mixed_opd_combines_per_sample_task_advantage(self): + first = self._scored_state( + 1, + behavior_logprobs=[-1.0], + teacher_logprobs=[-0.75], + reward=10.0, + ) + second = self._scored_state( + 2, + behavior_logprobs=[-1.0], + teacher_logprobs=[-1.25], + reward=10.0, + ) + third = self._scored_state( + 3, + behavior_logprobs=[-1.0], + teacher_logprobs=[-0.5], + reward=-1.0, + ) + estimator = _FixedAdvantageEstimator([2.0, 2.0, -2.0]) + + advantages = compute_pg_opd_token_advantages( + [first, second, third], + config=_algorithm_config(task_adv_weight=0.5, opd_adv_weight=2.0), + task_adv_estimator=estimator, + ) + + torch.testing.assert_close(torch.cat(advantages), torch.tensor([1.5, 0.5, 0.0])) + self.assertEqual(estimator.calls[0][0].tolist(), [10.0, 10.0, -1.0]) + self.assertEqual(estimator.calls[0][1], [first, second, third]) + + def test_mixed_opd_requires_reward(self): + state = self._scored_state( + 1, + behavior_logprobs=[-1.0], + teacher_logprobs=[-0.5], + ) + config = _algorithm_config(task_adv_weight=1.0) + + with self.assertRaisesRegex(ValueError, "Reward score is required"): + compute_pg_opd_token_advantages( + [state], + config=config, + task_adv_estimator=_FixedAdvantageEstimator([1.0]), + ) + + def test_token_advantages_survive_packing_and_define_denominator(self): + state = self._scored_state( + 1, + behavior_logprobs=[-1.0, -2.0, -3.0], + teacher_logprobs=[-0.5, -2.5, -2.0], + response_mask=[1, 0, 1], + ) + response_advantages = compute_pg_opd_token_advantages( + [state], + config=_algorithm_config(), + task_adv_estimator=None, + )[0] + data_batches = [ + { + "seq_ctx": SequenceContext.from_input_ids((torch.tensor([[10, 11, 20, 21]]),), device="cpu"), + "shifted_labels": torch.tensor([[-100, 100, -100, 102]]), + "advantage": [0.0] + response_advantages.tolist(), + "rollout_logprobs": torch.tensor([[0.0, -1.0, -2.0, -3.0]]), + } + ] + + controller = TrainingController.__new__(TrainingController) + packed = controller._packing(data_batches, pack_max_length=6, language_cfg=None) + + self.assertEqual(packed[0]["shifted_labels"].shape, packed[0]["advantages"].shape) + self.assertEqual(packed[0]["advantages"].tolist(), [[0.0, 0.5, 0.0, 1.0, -100.0, -100.0]]) + + loss_config = GRPOLossConfig( + policy_loss_cfg={"loss_type": "vanilla", "cliprange_low": 0.2, "cliprange_high": 0.2} + ) + loss_context = loss_config.build( + { + "shifted_labels": packed[0]["shifted_labels"], + "advantages": packed[0]["advantages"], + "old_logprobs": torch.zeros_like(packed[0]["advantages"]), + } + ) + assert isinstance(loss_context, GRPOLossContext) + GRPOLossContext.build_batches([loss_context]) + torch.testing.assert_close( + loss_context.loss_kwargs.policy_loss_weight.cpu(), + torch.tensor([[0.0, 0.5, 0.0, 0.5, 0.0, 0.0]]), + ) + + def test_teacher_signal_changes_student_gradient(self): + def gradient(teacher_logprob: float) -> torch.Tensor: + state = self._scored_state( + 1, + behavior_logprobs=[-1.0], + teacher_logprobs=[teacher_logprob], + ) + advantage = compute_pg_opd_token_advantages( + [state], + config=_algorithm_config(), + task_adv_estimator=None, + )[0].unsqueeze(0) + current_logprob = torch.tensor([[-1.0]], requires_grad=True) + old_logprob = torch.tensor([[-1.0]]) + loss_config = GRPOLossConfig( + policy_loss_cfg={"loss_type": "vanilla", "cliprange_low": 0.2, "cliprange_high": 0.2} + ) + context = loss_config.build( + { + "shifted_labels": torch.tensor([[100]]), + "advantages": advantage, + "old_logprobs": old_logprob, + } + ) + assert isinstance(context, GRPOLossContext) + GRPOLossContext.build_batches([context]) + device = context.loss_kwargs.advantages.device + current_logprob = current_logprob.to(device).detach().requires_grad_() + loss = context.policy_loss_fn( + current_logprob, + context.loss_kwargs.old_logprobs, + context.loss_kwargs.advantages, + context.loss_kwargs.policy_loss_weight, + loss_config.policy_loss_cfg, + ) + loss.backward() + return current_logprob.grad.detach().cpu().clone() + + positive_signal_gradient = gradient(-0.5) + stronger_signal_gradient = gradient(0.0) + + self.assertFalse(torch.equal(positive_signal_gradient, stronger_signal_gradient)) + + def test_pure_and_mixed_opd_support_rollout_is_off_and_on(self): + for task_adv_weight in (0.0, 1.0): + for enable_is in (False, True): + with self.subTest(task_adv_weight=task_adv_weight, enable_is=enable_is): + state = self._scored_state( + 1, + behavior_logprobs=[-1.0, -1.0], + teacher_logprobs=[0.0, 0.0], + reward=1.0 if task_adv_weight else None, + ) + estimator = _FixedAdvantageEstimator([1.0]) if task_adv_weight else None + advantages = compute_pg_opd_token_advantages( + [state], + config=_algorithm_config(task_adv_weight=task_adv_weight), + task_adv_estimator=estimator, + )[0].unsqueeze(0) + rollout_is = ( + RolloutImportanceSampling( + rollout_is_mode="mask", + rollout_is_threshold=(2.0, 0.5), + rollout_is_mask_threshold=(2.0, 0.5), + ) + if enable_is + else RolloutImportanceSampling() + ) + loss_config = GRPOLossConfig( + policy_loss_cfg={"loss_type": "vanilla", "cliprange_low": 0.2, "cliprange_high": 0.2}, + rollout_is=rollout_is, + ) + context = loss_config.build( + { + "shifted_labels": torch.tensor([[100, 101]]), + "advantages": advantages, + "rollout_logprobs": torch.tensor([[-1.0, -1.0]]), + "old_logprobs": torch.tensor([[-1.0, -1.0 + math.log(10.0)]]), + } + ) + assert isinstance(context, GRPOLossContext) + context.compute_rollout_is(None, torch.tensor([2])) # type: ignore[arg-type] + GRPOLossContext.build_batches([context]) + + device = context.loss_kwargs.advantages.device + current_logprobs = torch.tensor([[-1.0, -1.0]], device=device, requires_grad=True) + loss = context.policy_loss_fn( + current_logprobs, + context.loss_kwargs.old_logprobs, + context.loss_kwargs.advantages, + context.loss_kwargs.policy_loss_weight, + loss_config.policy_loss_cfg, + ) + loss.backward() + + if enable_is: + self.assertIsNotNone(context.loss_kwargs.is_weights) + self.assertEqual(context.loss_kwargs.shifted_labels.tolist(), [[100, -100]]) + self.assertEqual(current_logprobs.grad[0, 1].item(), 0.0) + else: + self.assertIsNone(context.loss_kwargs.is_weights) + self.assertEqual(context.loss_kwargs.shifted_labels.tolist(), [[100, 101]]) + self.assertNotEqual(current_logprobs.grad[0, 1].item(), 0.0) + + class TestOPDConfig(unittest.TestCase): def test_rejects_duplicate_or_unknown_teacher_names(self): teacher = OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:1") diff --git a/tests/rl/test_prepare_train_data.py b/tests/rl/test_prepare_train_data.py index 029a3f0306..81a364c889 100644 --- a/tests/rl/test_prepare_train_data.py +++ b/tests/rl/test_prepare_train_data.py @@ -19,6 +19,7 @@ import torch from xtuner.v1.data_proto.rl_data import RolloutState, Status +from xtuner.v1.rl.on_policy_distillation import OPDConfig, OPDTeacherConfig from xtuner.v1.train.rl_trainer import BaseRLTrainer @@ -33,13 +34,23 @@ def compute(self, rewards_tensor, group): class TestPrepareTrainData(unittest.TestCase): - def _build_trainer(self, advantages: list[float]): + def _build_trainer(self, advantages: list[float], opd_config: OPDConfig | None = None): trainer = BaseRLTrainer.__new__(BaseRLTrainer) trainer._advantage_estimator = _FakeAdvantageEstimator(advantages) + trainer._opd_config = opd_config trainer.tokenizer = MagicMock(return_value={"input_ids": torch.tensor([[999]])}) trainer.logger = MagicMock() return trainer + @staticmethod + def _opd_config(*, task_adv_weight: float = 0.0, opd_adv_weight: float = 1.0) -> OPDConfig: + return OPDConfig( + task_adv_weight=task_adv_weight, + opd_adv_weight=opd_adv_weight, + teachers=[OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:1")], + data_source_teacher_map={"math": "teacher"}, + ) + def _state( self, *, @@ -130,6 +141,68 @@ def test_multi_sample_group_uses_each_sample_reward_and_advantage(self): self.assertEqual(info["advantages/max"], 1.5) self.assertEqual(trainer._advantage_estimator.calls[0][0].tolist(), [3.0, -1.0]) + def test_pure_opd_builds_aligned_token_advantages_without_reward(self): + trainer = self._build_trainer([99.0], self._opd_config()) + state = self._state( + prompt_ids=[10, 11], + response_ids=[20, 21, 2], + logprobs=[-1.0, -2.0, -3.0], + response_mask=[1, 0, 1], + ) + state.reward = None + state.teacher_tokens = [20, 21, 2] + state.teacher_logprobs = [-0.5, -2.5, -2.0] + + data_batches, info = self._prepare(trainer, [[state]]) + + self.assertEqual(data_batches[0]["seq_ctx"].input_ids.tolist(), [[10, 11, 20, 21]]) + self.assertEqual(data_batches[0]["shifted_labels"].tolist(), [[-100, 20, -100, 2]]) + self.assertEqual(data_batches[0]["advantage"], [0.0, 0.5, 0.0, 1.0]) + self.assertEqual(len(data_batches[0]["advantage"]), data_batches[0]["shifted_labels"].numel()) + self.assertEqual(trainer._advantage_estimator.calls, []) + self.assertEqual(info["training_samples"], 1) + self.assertEqual(info["rewards/mean"], 0.0) + + def test_mixed_opd_combines_cluster_task_and_token_advantages(self): + trainer = self._build_trainer( + [2.0, -1.0], + self._opd_config(task_adv_weight=0.5, opd_adv_weight=2.0), + ) + first = self._state( + uid=1, + response_ids=[20, 2], + logprobs=[-1.0, -1.0], + reward={"score": 3.0}, + ) + first.teacher_tokens = [20, 2] + first.teacher_logprobs = [-0.5, -1.0] + second = self._state( + uid=2, + response_ids=[30, 2], + logprobs=[-1.0, -1.0], + reward={"score": -1.0}, + ) + second.teacher_tokens = [30, 2] + second.teacher_logprobs = [-1.5, -0.5] + + data_batches, info = self._prepare(trainer, [[first, second]]) + + self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 2.0, 1.0]) + self.assertEqual(data_batches[1]["advantage"], [0.0, 0.0, -1.5, 0.5]) + self.assertEqual(trainer._advantage_estimator.calls[0][0].tolist(), [3.0, -1.0]) + self.assertEqual(info["rewards/min"], -1.0) + self.assertEqual(info["rewards/max"], 3.0) + + def test_mixed_opd_missing_reward_fails_fast(self): + trainer = self._build_trainer([1.0], self._opd_config(task_adv_weight=1.0)) + state = self._state(response_ids=[20], logprobs=[-1.0]) + state.reward = None + state.teacher_tokens = [20] + state.teacher_logprobs = [-0.5] + + with self.assertRaisesRegex(ValueError, "Reward score is required"): + self._prepare(trainer, [[state]]) + 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]) diff --git a/xtuner/v1/rl/on_policy_distillation.py b/xtuner/v1/rl/on_policy_distillation.py index 69f6118dfa..8f49ca05fd 100644 --- a/xtuner/v1/rl/on_policy_distillation.py +++ b/xtuner/v1/rl/on_policy_distillation.py @@ -6,9 +6,11 @@ from typing import Literal, cast import httpx +import torch from pydantic import BaseModel, ConfigDict, Field, model_validator from xtuner.v1.data_proto.rl_data import RolloutState, Status +from xtuner.v1.rl.advantage import AdvantageEstimator class OPDTeacherConfig(BaseModel): @@ -127,3 +129,54 @@ def route_teacher_client( data_source = state.extra_fields["origin_data_source"] teacher_name = data_source_teacher_map[data_source] return teacher_clients[teacher_name] + + +def compute_pg_opd_token_advantages( + group: list[RolloutState], + *, + config: OPDConfig, + task_adv_estimator: AdvantageEstimator | None, +) -> list[torch.Tensor]: + opd_advantages: list[torch.Tensor] = [] + response_masks: list[torch.Tensor] = [] + for state in group: + response_ids = cast(list[int], state.response_ids) + behavior_logprobs = cast(list[float], state.logprobs) + teacher_logprobs = cast(list[float], state.teacher_logprobs) + + behavior_logprobs_t = torch.tensor(behavior_logprobs, dtype=torch.float32) + teacher_logprobs_t = torch.tensor(teacher_logprobs, dtype=torch.float32) + response_mask = state.response_mask + if not response_mask: + response_mask_t = torch.ones(len(response_ids), dtype=torch.float32) + else: + response_mask_t = torch.tensor(response_mask, dtype=torch.float32) + opd_advantages.append(teacher_logprobs_t - behavior_logprobs_t) + response_masks.append(response_mask_t) + + task_advantages = [0.0] * len(group) + if config.task_adv_weight > 0: + rewards: list[float] = [] + for state in group: + if state.reward is None or "score" not in state.reward: + raise ValueError(f"Reward score is required for mixed PG-OPD rollout {state.rollout_id}") + rewards.append(float(state.reward["score"])) + + task_advantages = ( + cast(AdvantageEstimator, task_adv_estimator) + .compute( + torch.tensor(rewards, dtype=torch.float32), + group, + ) + .tolist() + ) + + return [ + (config.task_adv_weight * task_advantage + config.opd_adv_weight * opd_advantage) * response_mask + for task_advantage, opd_advantage, response_mask in zip( + task_advantages, + opd_advantages, + response_masks, + strict=True, + ) + ] diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 9e810ad162..28e8d15253 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -33,6 +33,7 @@ ) from xtuner.v1.rl.agent_loop_manager.produce_utils import default_should_continue_fn from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.on_policy_distillation import OPDConfig, compute_pg_opd_token_advantages from xtuner.v1.rl.replay_buffer import ( AsyncReplayBufferConfig, SyncReplayBufferConfig, @@ -342,6 +343,7 @@ class BaseRLTrainerConfig(BaseModel): total_epochs: int | None = None train_batch_size: int advantage_estimator_config: BaseAdvantageConfig = Field(default_factory=GRPOAdvantageConfig) + opd_config: OPDConfig | None = None sync_weights_interval: int = 1 enable_evaluate: bool = True @@ -594,6 +596,7 @@ def _init_common(self, cfg: BaseRLTrainerConfig, *, meta_path: str, logger_tag: self._init_rollout_config(cfg, log_dir) self._ensure_rollout_proxy_config(cfg) self._init_runtime_flags(cfg) + self._opd_config = cfg.opd_config self._advantage_estimator = cfg.advantage_estimator_config.build() self._cpu_resource_manager: CPUResourceManager | None = None self._num_workers = 1.0 @@ -1014,13 +1017,14 @@ def _prepare_train_data( rewards_list = [] # 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. rewards_list stays per-segment for counts. + # rewards/* from being weighted by segment count. cluster_rewards_list: list[float] = [] advantages_list = [] prompt_len_list = [] response_len_list = [] tool_turns_list: list[int] = [] training_tokens = 0 + training_samples = 0 data_batches = [] @@ -1028,7 +1032,10 @@ def _prepare_train_data( 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) + # When opd_config is not None, PG-OPD currently supports only single-turn rollouts. + opd_config = self._opd_config 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 @@ -1040,39 +1047,57 @@ def _prepare_train_data( assert prompt_ids is not None and len(prompt_ids) > 0, ( f"Prompt ids cannot be None or empty in data: {group[0]}" ) - rewards = [] - # Agentic rollouts may split one model session into multiple trainable segments. - # Compute the group advantage once per session, then broadcast it back to each segment. - cluster_index_by_key: dict[Any, int] = {} - cluster_rewards: list[float] = [] - cluster_representatives: list[RolloutState] = [] - sample_cluster_indices: list[int] = [] for data in group: - assert data.reward is not None and "score" in data.reward, ( - f"Reward is missing or does not contain 'score' key in data: {data}" - ) - reward = float(data.reward["score"]) - rewards.append(reward) - # 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) # 有可能有重复,但是没有其他更好办法 turns = data.extra_fields.get("agent_tool_turns") if isinstance(turns, int): tool_turns_list.append(turns) - rewards_list.extend(rewards) - cluster_rewards_list.extend(cluster_rewards) - 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] + opd_token_advantages: list[torch.Tensor] | None = None + if opd_config is not None: + opd_token_advantages = compute_pg_opd_token_advantages( + group, + config=opd_config, + task_adv_estimator=self._advantage_estimator, + ) + sample_advantages: list[float] = [] + + if opd_config.task_adv_weight > 0: + rewards = [float(cast(dict[str, Any], data.reward)["score"]) for data in group] + rewards_list.extend(rewards) + cluster_rewards_list.extend(rewards) + else: + rewards = [] + # Agentic rollouts may split one model session into multiple trainable segments. + # Compute the group advantage once per session, then broadcast it back to each segment. + cluster_index_by_key: dict[Any, int] = {} + cluster_rewards: list[float] = [] + cluster_representatives: list[RolloutState] = [] + sample_cluster_indices: list[int] = [] + for data in group: + assert data.reward is not None and "score" in data.reward, ( + f"Reward is missing or does not contain 'score' key in data: {data}" + ) + reward = float(data.reward["score"]) + rewards.append(reward) + # 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) + + rewards_list.extend(rewards) + cluster_rewards_list.extend(cluster_rewards) + 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): @@ -1177,12 +1202,15 @@ def _prepare_train_data( shifted_labels = [-100] * (len(prompt_ids) - 1) + response_labels shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) - # 根据 response_mask 计算新的 advantages - advatnages_val = sample_advantages[i] - actual_advantages = [advatnages_val] * len(prompt_ids) + [ - 0.0 if mask == 0 else advatnages_val for mask in response_mask - ] - advantages_list.extend(actual_advantages[:-1]) + if opd_token_advantages is None: + advatnages_val = sample_advantages[i] + actual_advantages = [advatnages_val] * len(prompt_ids) + [ + 0.0 if mask == 0 else advatnages_val for mask in response_mask + ] + advantages_list.extend(actual_advantages[:-1]) + else: + actual_advantages = [0.0] * (len(prompt_ids) - 1) + opd_token_advantages[i].tolist() + advantages_list.extend(actual_advantages) assert len(input_ids) <= pack_max_length, f"{len(input_ids)} vs {pack_max_length}" training_tokens += len(input_ids) @@ -1214,8 +1242,8 @@ def _prepare_train_data( if not XTUNER_DETERMINISTIC: random.shuffle(data_batches) - # rewards/* report the per-session reward distribution; batch_size/training_samples below - # still use rewards_list (per-segment) so counts reflect the actual training samples. + # 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() @@ -1223,8 +1251,8 @@ def _prepare_train_data( raw_rewards_mean = raw_rewards_sum / raw_rewards_count if raw_rewards_count > 0 else rewards_t.mean().item() info_dict = { - "batch_size": len(rewards_list), - "training_samples": len(rewards_list), + "batch_size": training_samples, + "training_samples": training_samples, "training_tokens": training_tokens, "rewards/mean": rewards_t.mean().item(), "rewards/min": rewards_t.min().item(), From 50ca5953fa5e80e49edc87930d9099b2ab38f197 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Wed, 22 Jul 2026 08:07:46 +0000 Subject: [PATCH 4/8] [Feat] Propagate PG-OPD config to training agent loops --- xtuner/v1/rl/agent_loop/agent_loop.py | 12 ++------ .../agent_loop_manager/agent_loop_manager.py | 3 ++ .../disagg_agent_loop_manager.py | 3 ++ xtuner/v1/rl/on_policy_distillation.py | 30 +++++++++++++++---- xtuner/v1/train/rl_trainer.py | 7 ++++- 5 files changed, 38 insertions(+), 17 deletions(-) diff --git a/xtuner/v1/rl/agent_loop/agent_loop.py b/xtuner/v1/rl/agent_loop/agent_loop.py index 754087af96..4fa21fdf0f 100644 --- a/xtuner/v1/rl/agent_loop/agent_loop.py +++ b/xtuner/v1/rl/agent_loop/agent_loop.py @@ -16,6 +16,7 @@ OPDConfig, TeacherLogprobClient, route_teacher_client, + validate_opd_sample_params, ) from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.rl.rollout.constants import AGENT_LOOP_RAY_GENERATE_MAX_CONCURRENCY @@ -214,6 +215,7 @@ def configure_opd(self, opd_config: OPDConfig | None) -> None: if opd_config is None: return + validate_opd_sample_params(self.sample_params) self.teacher_clients = {teacher.name: TeacherLogprobClient(teacher) for teacher in opd_config.teachers} self.data_source_teacher_map = dict(opd_config.data_source_teacher_map) @@ -325,9 +327,6 @@ async def pause(self) -> None: finally: self._judger_pause_event.clear() - async def close(self) -> None: - await asyncio.gather(*(client.close() for client in self.teacher_clients.values())) - class RouterAgentLoop: def __init__(self, workers: list[RayAgentLoopProxy], rollout_ctl: RolloutController): @@ -379,9 +378,6 @@ async def pause(self) -> None: *(worker.pause.remote() for worker in self.workers), ) - async def close(self) -> None: - await asyncio.gather(*(worker.close.remote() for worker in self.workers)) - async def get_agent_loop_rollout_ctl(agent_loop: AgentLoopSpec) -> RolloutController: rollout_ctl = getattr(agent_loop, "rollout_ctl", None) @@ -431,10 +427,6 @@ async def get_rollout_ctl(self): async def pause(self) -> None: return await self.agent_loop.pause() - @ray_method - async def close(self) -> None: - return await self.agent_loop.close() - RayAgentLoop = cast( ActorClass[AgentLoopActor], diff --git a/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py b/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py index 3aa0b0e04c..57f0e92f4d 100644 --- a/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py +++ b/xtuner/v1/rl/agent_loop_manager/agent_loop_manager.py @@ -10,6 +10,7 @@ from xtuner.v1.data_proto.rl_data import Status from xtuner.v1.rl.agent_loop import AgentLoopConfig from xtuner.v1.rl.judger import ComposedJudgerConfig, JudgerConfig, build_judger +from xtuner.v1.rl.on_policy_distillation import OPDConfig from xtuner.v1.rl.replay_buffer import ReplayBuffer from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.utils import get_logger @@ -130,6 +131,7 @@ def build( replay_buffer: ReplayBuffer, logger=None, sync_weights_interval: int = 1, + opd_config: OPDConfig | None = None, ) -> "AgentLoopManager": tasks = self.tasks if isinstance(self.tasks, list) else [self.tasks] if not tasks: @@ -146,6 +148,7 @@ def build( rollout_controller=rollout_controller, judger=build_judger(task_cfg.judger_config) if task_cfg.judger_config is not None else None, logger=logger, + opd_config=opd_config, ) produce_strategy = task_cfg.produce_strategy_config.build( sync_weights_interval=sync_weights_interval, diff --git a/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py b/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py index ef2ff03434..e070b0026c 100644 --- a/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py +++ b/xtuner/v1/rl/agent_loop_manager/disagg_agent_loop_manager.py @@ -11,6 +11,7 @@ from xtuner.v1.data_proto.rl_data import Status from xtuner.v1.rl.agent_loop import AgentLoopConfig from xtuner.v1.rl.judger import ComposedJudgerConfig, JudgerConfig, build_judger +from xtuner.v1.rl.on_policy_distillation import OPDConfig from xtuner.v1.rl.replay_buffer import ReplayBuffer from xtuner.v1.rl.rollout import RolloutController from xtuner.v1.utils import get_logger @@ -70,6 +71,7 @@ def build( replay_buffer: ReplayBuffer, logger=None, sync_weights_interval: int = 1, + opd_config: OPDConfig | None = None, ) -> "DisaggAgentLoopManager": tasks = self.tasks if isinstance(self.tasks, list) else [self.tasks] if not tasks: @@ -86,6 +88,7 @@ def build( rollout_controller=rollout_controller, judger=build_judger(task_cfg.judger_config) if task_cfg.judger_config is not None else None, logger=logger, + opd_config=opd_config, ) produce_strategy = task_cfg.produce_strategy_config.build( sync_weights_interval=sync_weights_interval, diff --git a/xtuner/v1/rl/on_policy_distillation.py b/xtuner/v1/rl/on_policy_distillation.py index 8f49ca05fd..739495ea32 100644 --- a/xtuner/v1/rl/on_policy_distillation.py +++ b/xtuner/v1/rl/on_policy_distillation.py @@ -3,13 +3,13 @@ import asyncio import math import time -from typing import Literal, cast +from typing import Any, Literal, cast import httpx import torch from pydantic import BaseModel, ConfigDict, Field, model_validator -from xtuner.v1.data_proto.rl_data import RolloutState, Status +from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status from xtuner.v1.rl.advantage import AdvantageEstimator @@ -44,8 +44,29 @@ def validate_teacher_names(self) -> OPDConfig: return self +def validate_opd_sample_params(sample_params: SampleParams) -> None: + identity_sampling_params: dict[str, Any] = { + "temperature": 1.0, + "top_p": 1.0, + "top_k": 0, + "repetition_penalty": 1.0, + "presence_penalty": 0.0, + "frequency_penalty": 0.0, + "min_tokens": 0, + } + non_identity_params = { + name: getattr(sample_params, name) + for name, expected in identity_sampling_params.items() + if getattr(sample_params, name) != expected + } + if non_identity_params: + raise ValueError(f"PG-OPD requires identity student sampling, got {non_identity_params}") + if not sample_params.return_logprob or not sample_params.return_token_ids: + raise ValueError("PG-OPD requires return_logprob=True and return_token_ids=True") + + class TeacherLogprobClient: - """Minimal asynchronous client for one external SGLang teacher.""" + """Asynchronous SGLang teacher client scoped to one AgentLoop.""" def __init__(self, config: OPDTeacherConfig) -> None: self.config = config @@ -116,9 +137,6 @@ def _parse_response( raise ValueError("Teacher logprobs contain NaN or Inf") return teacher_tokens, teacher_logprobs - async def close(self) -> None: - await self._client.aclose() - def route_teacher_client( state: RolloutState, diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 28e8d15253..e650ad1ca3 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -33,7 +33,10 @@ ) from xtuner.v1.rl.agent_loop_manager.produce_utils import default_should_continue_fn from xtuner.v1.rl.evaluator import EvaluatorConfig -from xtuner.v1.rl.on_policy_distillation import OPDConfig, compute_pg_opd_token_advantages +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + compute_pg_opd_token_advantages, +) from xtuner.v1.rl.replay_buffer import ( AsyncReplayBufferConfig, SyncReplayBufferConfig, @@ -711,6 +714,7 @@ def _build_agent_loop_components(self, cfg: BaseRLTrainerConfig, replay_buffer) replay_buffer=replay_buffer, logger=self.logger, sync_weights_interval=cfg.sync_weights_interval, + opd_config=self._opd_config, ) self.agent_loop_manager = cast(AgentLoopManager | DisaggAgentLoopManager, agent_loop_manager) @@ -725,6 +729,7 @@ def _build_agent_loop_components(self, cfg: BaseRLTrainerConfig, replay_buffer) replay_buffer=replay_buffer, logger=self.logger, sync_weights_interval=cfg.sync_weights_interval, + opd_config=None, ), ) From f2eff6b2ca54b6e6e309106f4c22d7eee316291b Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Mon, 27 Jul 2026 09:38:36 +0000 Subject: [PATCH 5/8] [Fix] use old_logprobs instead of rollout_olgprobs to compute reverse kl --- .../rl_dapo_math_opd.py | 170 +++ recipe/on_policy_distillation/run_pg_opd.sh | 147 +++ .../start_pg_opd_teacher.sh | 19 + tests/rl/test_on_policy_distillation.py | 1163 +++++++---------- tests/rl/test_prepare_train_data.py | 87 +- xtuner/v1/rl/loss/base_loss.py | 9 + xtuner/v1/rl/on_policy_distillation.py | 77 +- xtuner/v1/rl/trainer/controller.py | 24 + xtuner/v1/rl/trainer/worker.py | 22 + xtuner/v1/train/rl_trainer.py | 43 +- 10 files changed, 941 insertions(+), 820 deletions(-) create mode 100644 recipe/on_policy_distillation/rl_dapo_math_opd.py create mode 100755 recipe/on_policy_distillation/run_pg_opd.sh create mode 100755 recipe/on_policy_distillation/start_pg_opd_teacher.sh diff --git a/recipe/on_policy_distillation/rl_dapo_math_opd.py b/recipe/on_policy_distillation/rl_dapo_math_opd.py new file mode 100644 index 0000000000..8e41b60439 --- /dev/null +++ b/recipe/on_policy_distillation/rl_dapo_math_opd.py @@ -0,0 +1,170 @@ +import os +from pathlib import Path + +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.data_proto.rl_data import SampleParams +from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig +from xtuner.v1.datasets.rl_tokenize_fn import RLTextTokenizeFnConfig +from xtuner.v1.model import get_model_config_from_hf +from xtuner.v1.rl.advantage import GRPOAdvantageConfig +from xtuner.v1.rl.agent_loop import SingleTurnAgentLoopConfig +from xtuner.v1.rl.agent_loop_manager import ( + AgentLoopManagerConfig, + SamplerConfig, + SyncProduceStrategyConfig, + TaskSpecConfig, +) +from xtuner.v1.rl.loss import GRPOLossConfig +from xtuner.v1.rl.on_policy_distillation import OPDConfig, OPDTeacherConfig +from xtuner.v1.rl.replay_buffer import SyncReplayBufferConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.trainer import WorkerConfig +from xtuner.v1.rl.utils import AcceleratorResourcesConfig +from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig + + +work_dir = os.environ["WORK_DIR"] +model_path = os.environ["MODEL_PATH"] +data_path = os.environ["DATA_PATH"] +NNODE = int(os.environ.get("WORLD_SIZE", "1")) + +# basic settings +experimental_name = "dapo_math_opd" +total_train_steps = 300 +train_batch_size = 64 +prompt_repeat_k = 4 +rollout_tp_size = 1 +rollout_ep_size = 1 +max_prompt_length = 2048 +max_response_length = 16384 +pack_max_length = max_prompt_length + max_response_length +train_optimizer_steps = 1 + +# 1. resources +resources = AcceleratorResourcesConfig( + accelerator="GPU", + num_workers=4 * NNODE, + num_cpus_per_worker=12, + cpu_memory_per_worker=16 * 1024**3, # 16 GB +) + +# 2. rollout +rollout_config = RolloutConfig( + env=experimental_name, + device=resources.accelerator, + model_path=model_path, + dtype="bfloat16", + tensor_parallel_size=rollout_tp_size, + expert_parallel_size=rollout_ep_size, + gpu_memory_utilization=0.4, + context_length=max_response_length + max_prompt_length, + enable_return_routed_experts=False, + rollout_max_batch_size_per_instance=2048, +) + +# 3. train worker +lr_cfg = LRConfig(lr_type="constant", warmup_ratio=0, lr_min=1e-6) +fsdp_cfg = FSDPConfig( + torch_compile=False, + cpu_offload=False, + ep_size=1, + reduce_dtype="float32", +) +model_cfg = get_model_config_from_hf(Path(model_path)) +if hasattr(model_cfg, "balancing_loss_cfg"): + model_cfg.balancing_loss_cfg = None +if hasattr(model_cfg, "z_loss_cfg"): + model_cfg.z_loss_cfg = None +optim_cfg = AdamWConfig(lr=1e-6, foreach=False, weight_decay=0.1, betas=(0.9, 0.98)) +loss_cfg = GRPOLossConfig( + policy_loss_cfg={ + "cliprange_high": 0.28, + "cliprange_low": 0.2, + "loss_type": "vanilla", + "clip_ratio_c": 10.0, + "log_prob_diff_min": -20.0, + "log_prob_diff_max": 20.0, + }, + ignore_idx=-100, + use_kl_loss=False, + kl_loss_coef=0.0, + kl_loss_type="low_var_kl", + mode="chunk", + chunk_size=512, +) +train_worker_cfg = WorkerConfig( + model_cfg=model_cfg, + load_from=model_path, + optim_cfg=optim_cfg, + loss_cfg=loss_cfg, + lr_cfg=lr_cfg, + fsdp_cfg=fsdp_cfg, + sp_size=1, + optimizer_steps=train_optimizer_steps, + pack_max_length=pack_max_length, +) + +# 4. train agent loop manager +train_dataset = DatasetConfig(name=experimental_name, anno_path=data_path) +tokenizer_config = RLTextTokenizeFnConfig(max_length=max_prompt_length) +train_dataset_cfg = [{"dataset": train_dataset, "tokenize_fn": tokenizer_config}] +dataloader_cfg = DataloaderConfig( + dataset_config_list=train_dataset_cfg, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", +) +sampler_config = SamplerConfig( + dataloader_cfg=dataloader_cfg, + prompt_repeat_k=prompt_repeat_k, +) +training_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, +) +agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=training_sample_params, +) +produce_strategy_config = SyncProduceStrategyConfig() +agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="train_task", + agent_loop_config=agent_loop_config, + produce_strategy_config=produce_strategy_config, + sampler_config=sampler_config, + ), +) + +# 5. pure on-policy distillation +opd_config = OPDConfig( + mode="pg-opd", + task_adv_weight=0.0, + opd_adv_weight=1.0, + teachers=[OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:13141")], + data_source_teacher_map={"math_dapo": "teacher"}, +) + +trainer = RLColocateTrainerConfig( + resources=resources, + train_worker_cfg=train_worker_cfg, # TODO: uniform naming of cfg and config + rollout_config=rollout_config, + tokenizer_path=model_path, + replay_buffer_config=SyncReplayBufferConfig(), + agent_loop_manager_cfg=agent_loop_manager_cfg, + load_from=model_path, + train_batch_size=train_batch_size, + advantage_estimator_config=GRPOAdvantageConfig(eps=1e-8), + opd_config=opd_config, + enable_evaluate=False, + enable_initial_evaluate=False, + total_train_steps=total_train_steps, + work_dir=work_dir, + seed=1234, + debug_rollout=False, +) diff --git a/recipe/on_policy_distillation/run_pg_opd.sh b/recipe/on_policy_distillation/run_pg_opd.sh new file mode 100755 index 0000000000..b4f8aae7ee --- /dev/null +++ b/recipe/on_policy_distillation/run_pg_opd.sh @@ -0,0 +1,147 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../.." && pwd) + +if (( $# != 0 )); then + echo "This script does not accept positional arguments." >&2 + echo "Set STUDENT_MODEL_PATH, TEACHER_MODEL_PATH, and DATA_PATH before running it." >&2 + exit 2 +fi + +: "${STUDENT_MODEL_PATH:?STUDENT_MODEL_PATH is required}" +: "${TEACHER_MODEL_PATH:?TEACHER_MODEL_PATH is required}" +: "${DATA_PATH:?DATA_PATH is required}" + +STUDENT_CUDA_VISIBLE_DEVICES="0,1,2,3" +TEACHER_CUDA_VISIBLE_DEVICES="7" +TEACHER_HOST="127.0.0.1" +TEACHER_PORT="13141" +TEACHER_TP_SIZE="1" +TEACHER_CHUNKED_PREFILL_SIZE="4096" +TEACHER_GPU_MEMORY_UTILIZATION="0.6" +TEACHER_STARTUP_TIMEOUT_S="1200" + +WORK_DIR="${REPO_ROOT}/work_dirs/dapo_math_opd" +TEACHER_ENDPOINT="http://${TEACHER_HOST}:${TEACHER_PORT}" +TEACHER_LOG_FILE="${WORK_DIR}/teacher.log" + +if [[ ! -d "${STUDENT_MODEL_PATH}" ]]; then + echo "Student model directory does not exist: ${STUDENT_MODEL_PATH}" >&2 + exit 1 +fi +if [[ ! -d "${TEACHER_MODEL_PATH}" ]]; then + echo "Teacher model directory does not exist: ${TEACHER_MODEL_PATH}" >&2 + exit 1 +fi +if [[ ! -f "${DATA_PATH}" ]]; then + echo "Training data file does not exist: ${DATA_PATH}" >&2 + exit 1 +fi + +for required_command in python curl ray setsid; do + if ! command -v "${required_command}" >/dev/null 2>&1; then + echo "Required command is not available: ${required_command}" >&2 + exit 1 + fi +done + +python -c "import sglang" >/dev/null +mkdir -p "${WORK_DIR}" + +if curl -sf --max-time 2 "${TEACHER_ENDPOINT}/health_generate" >/dev/null; then + echo "A teacher service is already running at ${TEACHER_ENDPOINT}" >&2 + exit 1 +fi + +TEACHER_PID="" +TRAINING_STARTED=0 + +cleanup() { + local exit_code=$? + local attempt + + trap - EXIT INT TERM + + if [[ -n "${TEACHER_PID}" ]]; then + kill -TERM -- "-${TEACHER_PID}" 2>/dev/null || true + for ((attempt = 0; attempt < 30; attempt++)); do + if ! kill -0 -- "-${TEACHER_PID}" 2>/dev/null; then + break + fi + sleep 1 + done + if kill -0 -- "-${TEACHER_PID}" 2>/dev/null; then + kill -KILL -- "-${TEACHER_PID}" 2>/dev/null || true + fi + wait "${TEACHER_PID}" 2>/dev/null || true + fi + + if (( TRAINING_STARTED )); then + ray stop --force >/dev/null 2>&1 || true + fi + + exit "${exit_code}" +} + +wait_for_teacher() { + local deadline=$((SECONDS + TEACHER_STARTUP_TIMEOUT_S)) + + while (( SECONDS < deadline )); do + if ! kill -0 "${TEACHER_PID}" 2>/dev/null; then + echo "Teacher process exited before becoming ready." >&2 + tail -n 50 "${TEACHER_LOG_FILE}" >&2 || true + return 1 + fi + if curl -sf --max-time 2 "${TEACHER_ENDPOINT}/health_generate" >/dev/null; then + return 0 + fi + echo "Waiting for teacher service at ${TEACHER_ENDPOINT}..." + sleep 5 + done + + echo "Teacher service did not become ready within ${TEACHER_STARTUP_TIMEOUT_S} seconds." >&2 + tail -n 50 "${TEACHER_LOG_FILE}" >&2 || true + return 1 +} + +trap cleanup EXIT +trap "exit 130" INT +trap "exit 143" TERM + +echo "Starting teacher model: ${TEACHER_MODEL_PATH}" +echo "Teacher GPUs: ${TEACHER_CUDA_VISIBLE_DEVICES}" +echo "Teacher log: ${TEACHER_LOG_FILE}" + +setsid env \ + CUDA_VISIBLE_DEVICES="${TEACHER_CUDA_VISIBLE_DEVICES}" \ + PYTHONUNBUFFERED=1 \ + python -m sglang.launch_server \ + --model-path "${TEACHER_MODEL_PATH}" \ + --host "${TEACHER_HOST}" \ + --port "${TEACHER_PORT}" \ + --tp "${TEACHER_TP_SIZE}" \ + --chunked-prefill-size "${TEACHER_CHUNKED_PREFILL_SIZE}" \ + --mem-fraction-static "${TEACHER_GPU_MEMORY_UTILIZATION}" \ + >"${TEACHER_LOG_FILE}" 2>&1 & +TEACHER_PID=$! + +wait_for_teacher +curl -sS --max-time 10 "${TEACHER_ENDPOINT}/get_model_info" +echo +echo "Teacher service is ready at ${TEACHER_ENDPOINT}" +echo "Starting Pure PG-OPD training with student GPUs: ${STUDENT_CUDA_VISIBLE_DEVICES}" + +export WORK_DIR +export PYTHONUNBUFFERED=1 + +cd "${REPO_ROOT}" +TRAINING_STARTED=1 +CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES}" \ + bash -o pipefail examples/v1/scripts/run_rl.sh \ + recipe/on_policy_distillation/rl_dapo_math_opd.py \ + sglang \ + "${STUDENT_MODEL_PATH}" \ + "${DATA_PATH}" diff --git a/recipe/on_policy_distillation/start_pg_opd_teacher.sh b/recipe/on_policy_distillation/start_pg_opd_teacher.sh new file mode 100755 index 0000000000..575e0210f2 --- /dev/null +++ b/recipe/on_policy_distillation/start_pg_opd_teacher.sh @@ -0,0 +1,19 @@ +#!/usr/bin/env bash + +set -euo pipefail + +TEACHER_MODEL_PATH=${1:?"Usage: $0 TEACHER_MODEL_PATH"} +TEACHER_HOST=${TEACHER_HOST:-127.0.0.1} +TEACHER_PORT=${TEACHER_PORT:-13141} +TEACHER_TP_SIZE=${TEACHER_TP_SIZE:-1} +TEACHER_CHUNKED_PREFILL_SIZE=${TEACHER_CHUNKED_PREFILL_SIZE:-4096} +TEACHER_GPU_MEMORY_UTILIZATION=${TEACHER_GPU_MEMORY_UTILIZATION:-0.6} +PYTHON_EXECUTABLE=${PYTHON_EXECUTABLE:-python} + +exec "${PYTHON_EXECUTABLE}" -m sglang.launch_server \ + --model-path "${TEACHER_MODEL_PATH}" \ + --host "${TEACHER_HOST}" \ + --port "${TEACHER_PORT}" \ + --tp "${TEACHER_TP_SIZE}" \ + --chunked-prefill-size "${TEACHER_CHUNKED_PREFILL_SIZE}" \ + --mem-fraction-static "${TEACHER_GPU_MEMORY_UTILIZATION}" diff --git a/tests/rl/test_on_policy_distillation.py b/tests/rl/test_on_policy_distillation.py index 937c169be8..df87baa047 100644 --- a/tests/rl/test_on_policy_distillation.py +++ b/tests/rl/test_on_policy_distillation.py @@ -1,714 +1,557 @@ -import asyncio -import json import math -import socket -import threading +import os +import signal +import subprocess import time import unittest -from collections.abc import Callable -from dataclasses import dataclass -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from typing import Any -from unittest.mock import MagicMock +from pathlib import Path +from unittest.mock import patch +import httpx import torch -from xtuner.v1.data_proto.rl_data import ( - RolloutState, - SampleParams, - Status, - get_group_status, - reset_rollout_response, -) -from xtuner.v1.data_proto.sequence_context import SequenceContext -from xtuner.v1.rl.agent_loop.single_turn_agent_loop import SingleTurnAgentLoop -from xtuner.v1.rl.loss import GRPOLossConfig, GRPOLossContext +from xtuner.v1.data_proto.rl_data import RolloutState, Status +from xtuner.v1.rl.loss import GRPOLossConfig from xtuner.v1.rl.on_policy_distillation import ( OPDConfig, OPDTeacherConfig, TeacherLogprobClient, - compute_pg_opd_token_advantages, + apply_opd_kl_to_advantages, ) -from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig -from xtuner.v1.rl.rollout_is import RolloutImportanceSampling -from xtuner.v1.rl.trainer.controller import TrainingController - - -@dataclass -class _HTTPResponse: - body: Any - status: int = 200 - delay_s: float = 0.0 - - -class _EndpointState: - def __init__( - self, - responses: list[_HTTPResponse] | None = None, - on_request: Callable[[dict[str, Any]], None] | None = None, - ) -> None: - self.responses = list(responses or []) - self.on_request = on_request - self.requests: list[dict[str, Any]] = [] - self.headers: list[dict[str, str]] = [] - self.request_event = threading.Event() - self.max_active_requests = 0 - self._active_requests = 0 - self._lock = threading.Lock() - - def handle(self, payload: dict[str, Any], headers: dict[str, str]) -> _HTTPResponse: - with self._lock: - self.requests.append(payload) - self.headers.append(headers) - self._active_requests += 1 - self.max_active_requests = max(self.max_active_requests, self._active_requests) - response = self.responses.pop(0) if self.responses else self._success_response(payload) - self.request_event.set() - if self.on_request is not None: - self.on_request(payload) - try: - if response.delay_s: - time.sleep(response.delay_s) - return response - finally: - with self._lock: - self._active_requests -= 1 - - @staticmethod - def _success_response(payload: dict[str, Any]) -> _HTTPResponse: - scored_tokens = payload["input_ids"][payload["logprob_start_len"] :] - return _HTTPResponse( - body={ - "meta_info": { - "input_token_logprobs": [[-0.1 - index / 100, token] for index, token in enumerate(scored_tokens)] - } - } +from xtuner.v1.rl.utils import find_free_ports + +REPO_ROOT = Path(__file__).resolve().parents[2] +BASELINE_PATH = os.getenv("XTUNER_OPD_BASELINE") +TEACHER_MODEL_PATH = os.getenv("XTUNER_OPD_TEACHER_MODEL") +STUDENT_MODEL_PATH = os.getenv("XTUNER_OPD_STUDENT_MODEL") +TEACHER_STARTUP_TIMEOUT_S = float(os.getenv("XTUNER_OPD_TEACHER_STARTUP_TIMEOUT_S", "1200")) +TEACHER_START_SCRIPT = REPO_ROOT / "recipe/on_policy_distillation/start_pg_opd_teacher.sh" +TRAINER_CONFIG_PATH = REPO_ROOT / "recipe/on_policy_distillation/rl_dapo_math_opd.py" + + +def _wait_for_teacher(process: subprocess.Popen, endpoint: str) -> None: + deadline = time.monotonic() + TEACHER_STARTUP_TIMEOUT_S + with httpx.Client(timeout=1.0, trust_env=False) as client: + while time.monotonic() < deadline: + return_code = process.poll() + if return_code is not None: + raise RuntimeError(f"Teacher process exited during startup with code {return_code}") + try: + if client.get(f"{endpoint}/health_generate").status_code == 200: + return + except httpx.RequestError: + pass + time.sleep(1.0) + raise TimeoutError(f"Teacher did not become ready within {TEACHER_STARTUP_TIMEOUT_S} seconds") + + +def _build_debug_rollout_batch(samples: list[dict]) -> list[list[RolloutState]]: + train_batch = [] + for sample in samples: + sample_index = int(sample["sample_index"]) + group_index = sample.get("group_index") + prompt_ids = torch.as_tensor(sample["prompt_token_ids"], dtype=torch.long).tolist() + response_ids = torch.as_tensor(sample["response_token_ids"], dtype=torch.long).tolist() + rollout_logprobs = torch.as_tensor(sample["rollout_log_probs"], dtype=torch.float32).tolist() + teacher_logprobs = torch.as_tensor(sample["teacher_log_probs"], dtype=torch.float32).tolist() + response_mask = torch.as_tensor(sample["loss_mask"], dtype=torch.bool).int().tolist() + response = str(sample["response"]) + + train_batch.append( + [ + RolloutState( + rollout_id=sample_index, + group_id=sample_index if group_index is None else int(group_index), + message=[], + prompt_ids=prompt_ids, + tokens=prompt_ids, + response=response, + response_ids=response_ids, + logprobs=rollout_logprobs, + teacher_tokens=response_ids, + teacher_logprobs=teacher_logprobs, + response_mask=response_mask, + reward={"score": 0.0}, + finish_reason="stop", + status=Status.COMPLETED, + extra_fields={"origin_data_source": "baseline"}, + ) + ] ) + return train_batch -class _Handler(BaseHTTPRequestHandler): - def do_POST(self) -> None: - content_length = int(self.headers["Content-Length"]) - payload = json.loads(self.rfile.read(content_length)) - response = self.server.endpoint_state.handle(payload, dict(self.headers)) # type: ignore[attr-defined] - body = response.body if isinstance(response.body, bytes) else json.dumps(response.body).encode() - self.send_response(response.status) - self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(body))) - self.end_headers() - try: - self.wfile.write(body) - except (BrokenPipeError, ConnectionResetError): - pass - - def log_message(self, format: str, *args: Any) -> None: - return - - -class _FakeEndpoint: - def __init__( - self, - responses: list[_HTTPResponse] | None = None, - on_request: Callable[[dict[str, Any]], None] | None = None, - ) -> None: - self.state = _EndpointState(responses, on_request) - self.server = ThreadingHTTPServer(("127.0.0.1", 0), _Handler) - self.server.endpoint_state = self.state # type: ignore[attr-defined] - self.thread = threading.Thread(target=self.server.serve_forever, daemon=True) - self.thread.start() - - @property - def url(self) -> str: - host, port = self.server.server_address - return f"http://{host}:{port}" - - def close(self) -> None: - self.server.shutdown() - self.server.server_close() - self.thread.join(timeout=2) - - -def _state( - rollout_id: int, - *, - data_source: str = "math", - prompt_ids: list[int] | None = None, - response_ids: list[int] | None = None, -) -> RolloutState: - prompt_ids = prompt_ids or [100 + rollout_id] - response_ids = response_ids or [200 + rollout_id, 2] - return RolloutState( - rollout_id=rollout_id, - group_id=0, - message=[{"role": "user", "content": f"prompt {rollout_id}"}], - prompt_ids=prompt_ids, - tokens=prompt_ids, - response="response", - response_ids=response_ids, - logprobs=[-0.4] * len(response_ids), - response_mask=[1] * len(response_ids), - status=Status.COMPLETED, - extra_fields={"origin_data_source": data_source}, - ) - - -def _opd_config(endpoint_by_name: dict[str, str], data_source_teacher_map: dict[str, str], **kwargs) -> OPDConfig: - return OPDConfig( - teachers=[ - OPDTeacherConfig(name=name, endpoint=endpoint, **kwargs) for name, endpoint in endpoint_by_name.items() - ], - data_source_teacher_map=data_source_teacher_map, - ) - - -def _algorithm_config(*, task_adv_weight: float = 0.0, opd_adv_weight: float = 1.0) -> OPDConfig: - return OPDConfig( - task_adv_weight=task_adv_weight, - opd_adv_weight=opd_adv_weight, - teachers=[OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:1")], - data_source_teacher_map={"math": "teacher"}, - ) - - -class _FixedAdvantageEstimator: - def __init__(self, values: list[float]) -> None: - self.values = values - self.calls: list[tuple[torch.Tensor, list[RolloutState]]] = [] - - def compute(self, rewards: torch.Tensor, group: list[RolloutState]) -> torch.Tensor: - self.calls.append((rewards.clone(), group)) - return torch.tensor(self.values[: len(group)], dtype=torch.float32) +@unittest.skipUnless(BASELINE_PATH, "XTUNER_OPD_BASELINE is required") +class TestPGOPDAdvantage(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + super().setUpClass() + baseline = torch.load(BASELINE_PATH, map_location="cpu", weights_only=False) + cls.samples = baseline["samples"] - -class TestTeacherLogprobClient(unittest.IsolatedAsyncioTestCase): - def _endpoint(self, responses: list[_HTTPResponse] | None = None) -> _FakeEndpoint: - endpoint = _FakeEndpoint(responses) - self.addCleanup(endpoint.close) - return endpoint - - async def test_compute_logprobs_uses_sglang_prefill_protocol(self): - endpoint = self._endpoint() - client = TeacherLogprobClient( - OPDTeacherConfig(name="teacher", endpoint=endpoint.url, api_key="secret", max_retry_per_sample=0) + def test_compute_advantages_matches_baseline(self) -> None: + config = OPDConfig( + teachers=[OPDTeacherConfig(name="teacher", endpoint="http://unused")], + data_source_teacher_map={"baseline": "teacher"}, ) - self.addAsyncCleanup(client.close) - state = _state(1, prompt_ids=[10], response_ids=[20, 2]) - - result = await client.compute_logprobs(state) - - self.assertIs(result, state) - self.assertEqual(result.teacher_tokens, [20, 2]) - for actual, expected in zip(result.teacher_logprobs or [], [-0.11, -0.12], strict=True): - self.assertAlmostEqual(actual, expected) - payload = endpoint.state.requests[0] - self.assertEqual(payload["input_ids"], [10, 20, 2]) - self.assertEqual(payload["logprob_start_len"], 0) - self.assertEqual(payload["top_logprobs_num"], 0) - self.assertEqual( - payload["sampling_params"], - {"max_new_tokens": 0, "temperature": 1.0, "skip_special_tokens": False}, - ) - self.assertEqual(endpoint.state.headers[0]["Authorization"], "Bearer secret") - - async def test_invalid_teacher_responses_fail_without_signal(self): - responses = [ - _HTTPResponse(b"not-json"), - _HTTPResponse({"missing": "meta_info"}), - _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [-0.2, 20]]}}), - _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [-0.2, 20], [-0.3, 2], [-0.4, 3]]}}), - _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [-0.2, 999], [-0.3, 2]]}}), - _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [float("nan"), 20], [-0.3, 2]]}}), - _HTTPResponse({"meta_info": {"input_token_logprobs": [[-0.1, 10], [float("inf"), 20], [-0.3, 2]]}}), - ] - endpoint = self._endpoint(responses) - client = TeacherLogprobClient(OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=0)) - self.addAsyncCleanup(client.close) - - for rollout_id in range(len(responses)): - with self.subTest(rollout_id=rollout_id): - result = await client.compute_logprobs(_state(rollout_id, prompt_ids=[10], response_ids=[20, 2])) - self.assertEqual(result.status, Status.FAILED) - self.assertIsNone(result.teacher_tokens) - self.assertIsNone(result.teacher_logprobs) - self.assertIn("scoring failed after 1 attempts", result.error_msg or "") - - async def test_http_status_retries_then_succeeds(self): - endpoint = self._endpoint([_HTTPResponse({"error": "busy"}, status=503)]) - client = TeacherLogprobClient(OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=1)) - self.addAsyncCleanup(client.close) - - result = await client.compute_logprobs(_state(1)) - - self.assertEqual(result.status, Status.COMPLETED) - self.assertEqual(len(endpoint.state.requests), 2) - - async def test_http_status_retry_exhaustion_marks_failed(self): - for status in (400, 500): - with self.subTest(status=status): - endpoint = self._endpoint([_HTTPResponse({"error": "failed"}, status=status) for _ in range(2)]) - client = TeacherLogprobClient( - OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=1) + loss_cfg = GRPOLossConfig(policy_loss_cfg={"loss_type": "vanilla"}) + + for sample in self.samples: + with self.subTest(sample_index=sample["sample_index"]): + old_logprobs = torch.as_tensor(sample["old_log_probs"], dtype=torch.float32) + teacher_logprobs = torch.as_tensor(sample["teacher_log_probs"], dtype=torch.float32) + loss_mask = torch.as_tensor(sample["loss_mask"], dtype=torch.float32) + shifted_labels = torch.where( + loss_mask.bool(), + torch.zeros_like(loss_mask, dtype=torch.long), + torch.full_like(loss_mask, -100, dtype=torch.long), + ) + loss_ctx = loss_cfg.build( + { + "shifted_labels": shifted_labels, + "advantages": torch.zeros_like(old_logprobs), + "old_logprobs": old_logprobs, + "teacher_logprobs": teacher_logprobs, + } ) - self.addAsyncCleanup(client.close) + assert loss_ctx is not None + apply_opd_kl_to_advantages(loss_ctx, config=config) + actual_advantages = loss_ctx.loss_kwargs.advantages.cpu() + expected_advantages = torch.as_tensor(sample["advantages"], dtype=torch.float32) * loss_mask + try: + torch.testing.assert_close( + actual_advantages, + expected_advantages, + rtol=1e-5, + atol=1e-5, + ) + except AssertionError as error: + mismatch_indices = torch.nonzero( + ~torch.isclose(actual_advantages, expected_advantages, rtol=1e-5, atol=1e-5) + ).flatten() + mismatch_values = "\n".join( + ( + f"index={index}: " + f"actual={actual_advantages[index].item()!r}, " + f"expected={expected_advantages[index].item()!r}, " + f"abs_diff={abs(actual_advantages[index] - expected_advantages[index]).item()!r}" + ) + for index in mismatch_indices.tolist() + ) + raise AssertionError(f"{error}\n\nMismatched advantages:\n{mismatch_values}") from None + + +def _load_trainer_samples(capture_dir: Path) -> dict[int, dict]: + trainer_samples = {} + for capture_file in sorted(capture_dir.glob("rank_*.pt")): + for batch in torch.load(capture_file, map_location="cpu", weights_only=True): + shifted_labels = torch.as_tensor(batch["shifted_labels"], dtype=torch.long).reshape(-1) + old_log_probs = torch.as_tensor(batch["old_log_probs"], dtype=torch.float32).reshape(-1) + advantages = torch.as_tensor(batch["advantages"], dtype=torch.float32).reshape(-1) + boundaries = torch.as_tensor(batch["cu_seq_lens_q"], dtype=torch.long).tolist() + num_padding = int(batch["num_padding"]) + padding_start = shifted_labels.numel() - num_padding + + assert boundaries[len(batch["rollout_ids"])] == padding_start + torch.testing.assert_close( + shifted_labels[padding_start:], + torch.full_like(shifted_labels[padding_start:], -100), + rtol=0, + atol=0, + ) + + for rollout_id, start, end in zip(batch["rollout_ids"], boundaries[:-1], boundaries[1:]): + response_mask = shifted_labels[start:end] != -100 + trainer_samples[int(rollout_id)] = { + "response_token_ids": shifted_labels[start:end][response_mask], + "old_log_probs": old_log_probs[start:end][response_mask], + "advantages": advantages[start:end][response_mask], + } - result = await client.compute_logprobs(_state(status)) + return trainer_samples - self.assertEqual(result.status, Status.FAILED) - self.assertEqual(len(endpoint.state.requests), 2) - async def test_timeout_and_connection_error_use_bounded_retries(self): - endpoint = self._endpoint([_HTTPResponse({}, delay_s=0.05) for _ in range(2)]) - timeout_client = TeacherLogprobClient( - OPDTeacherConfig( - name="timeout", - endpoint=endpoint.url, - request_timeout_s=0.01, - max_retry_per_sample=1, +def _run_trainer_once( + baseline_samples: list[dict], + *, + baseline_path: Path, + student_model_path: Path, + debug_rollout_dir: Path, + capture_dir: Path, + run_dir: Path, +) -> dict[int, dict]: + import ray + + from xtuner.v1.rl.trainer.controller import TrainingController + from xtuner.v1.rl.trainer.worker import TrainingWorker, WorkerConfig + from xtuner.v1.train.rl_trainer import RLColocateTrainer + from xtuner.v1.utils import Config + + max_prompt_length = max(len(sample["prompt_token_ids"]) for sample in baseline_samples) + max_response_length = max(len(sample["response_token_ids"]) for sample in baseline_samples) + max_input_length = max(len(sample["token_ids"]) - 1 for sample in baseline_samples) + pack_max_length = math.ceil(max_input_length / 512) * 512 + num_workers = 8 + + class CapturingRLColocateTrainer(RLColocateTrainer): + def _prepare_train_data( + self, + data_groups, + pack_max_length, + raw_rewards_sum=0.0, + raw_rewards_count=0, + ): + data_batches, data_info = super()._prepare_train_data( + data_groups, + pack_max_length, + raw_rewards_sum=raw_rewards_sum, + raw_rewards_count=raw_rewards_count, ) - ) - self.addAsyncCleanup(timeout_client.close) - - timeout_result = await timeout_client.compute_logprobs(_state(1)) - - self.assertEqual(timeout_result.status, Status.FAILED) - self.assertEqual(len(endpoint.state.requests), 2) - - sock = socket.socket() - sock.bind(("127.0.0.1", 0)) - closed_port = sock.getsockname()[1] - sock.close() - connection_client = TeacherLogprobClient( - OPDTeacherConfig( - name="connection", - endpoint=f"http://127.0.0.1:{closed_port}", - request_timeout_s=0.1, - max_retry_per_sample=1, + rollout_states = (state for group in data_groups for state in group) + for data_batch, rollout_state in zip(data_batches, rollout_states): + data_batch["rollout_id"] = rollout_state.rollout_id + return data_batches, data_info + + class CapturingTrainingController(TrainingController): + def _packing(self, data_batches, pack_max_length, language_cfg): + pack_infos = self._get_pack_infos( + data_batches, + [data["seq_ctx"].input_ids.numel() for data in data_batches], + pack_max_length, ) - ) - self.addAsyncCleanup(connection_client.close) - - connection_result = await connection_client.compute_logprobs(_state(2)) - - self.assertEqual(connection_result.status, Status.FAILED) - - async def test_max_concurrency_limits_requests_per_client(self): - endpoint = self._endpoint([_HTTPResponse({}, delay_s=0.05) for _ in range(5)]) - client = TeacherLogprobClient( - OPDTeacherConfig(name="teacher", endpoint=endpoint.url, max_retry_per_sample=0, max_concurrency=2) - ) - self.addAsyncCleanup(client.close) - - results = await asyncio.gather(*(client.compute_logprobs(_state(index)) for index in range(5))) - - self.assertEqual(endpoint.state.max_active_requests, 2) - self.assertTrue(all(result.status == Status.FAILED for result in results)) - - -class TestAgentLoopTeacherScoring(unittest.IsolatedAsyncioTestCase): - def _endpoint( - self, - responses: list[_HTTPResponse] | None = None, - on_request: Callable[[dict[str, Any]], None] | None = None, - ) -> _FakeEndpoint: - endpoint = _FakeEndpoint(responses, on_request) - self.addCleanup(endpoint.close) - return endpoint - - def _loop(self, opd_config: OPDConfig) -> SingleTurnAgentLoop: - loop = SingleTurnAgentLoop.__new__(SingleTurnAgentLoop) - loop.rollout_ctl = MagicMock() - loop.sample_params = SampleParams(max_tokens=8) - loop.judger = None - loop.enable_batch_judge = False - loop._judger_pause_event = asyncio.Event() - loop.logger = MagicMock() - loop.configure_opd(opd_config) - self.addAsyncCleanup(loop.close) - return loop - - @staticmethod - def _complete(state: RolloutState) -> RolloutState: - state.status = Status.COMPLETED - return state - - async def test_eager_scoring_starts_before_other_samples_finish_generation(self): - endpoint = self._endpoint() - loop = self._loop(_opd_config({"teacher": endpoint.url}, {"math": "teacher"})) - release_second = asyncio.Event() - second_state = _state(2) - second_state.status = Status.INIT - - async def generate_sample(state, **kwargs): - if state is second_state: - await release_second.wait() - return self._complete(state) - - loop.generate_sample = generate_sample - task = asyncio.create_task(loop.collect_rollout_group([_state(1), second_state])) - - request_started = await asyncio.to_thread(endpoint.state.request_event.wait, 1.0) - self.assertTrue(request_started) - self.assertEqual(second_state.status, Status.INIT) - release_second.set() - result = await task - - self.assertEqual(get_group_status(result), Status.COMPLETED) - self.assertEqual(len(endpoint.state.requests), 2) - self.assertTrue(all(state.teacher_logprobs is not None for state in result)) - self.assertTrue(all("teacher_score_time_s" in state.extra_fields for state in result)) - - async def test_filter_false_skips_teacher_endpoint(self): - endpoint = self._endpoint() - loop = self._loop(_opd_config({"teacher": endpoint.url}, {"math": "teacher"})) - - async def generate_sample(state, **kwargs): - return self._complete(state) - - loop.generate_sample = generate_sample - result = await loop.collect_rollout_group( - [_state(1), _state(2)], - is_valid_sample_func=lambda group: False, - ) + packed_data_batches = super()._packing(data_batches, pack_max_length, language_cfg) + for packed_data, pack_info in zip(packed_data_batches, pack_infos): + packed_data["rollout_ids"] = [data_batches[index]["rollout_id"] for index in pack_info["indices"]] + return packed_data_batches - self.assertEqual(get_group_status(result), Status.FILTERED) - self.assertEqual(endpoint.state.requests, []) - self.assertTrue(all("teacher_score_time_s" not in state.extra_fields for state in result)) + class CapturingTrainingWorker(TrainingWorker): + trainer_capture_dir = capture_dir - async def test_filter_true_runs_before_lazy_teacher_scoring(self): - order = [] - endpoint = self._endpoint(on_request=lambda payload: order.append("teacher")) + def fit(self, data_batches, rollout_idx): + from unittest.mock import patch as mock_patch - def filter_func(group): - order.append("filter") - return True + from xtuner.v1.rl.trainer import worker as worker_module - loop = self._loop(_opd_config({"teacher": endpoint.url}, {"math": "teacher"})) + captured_batches = [ + { + "rollout_ids": data.get("rollout_ids", []), + "cu_seq_lens_q": data["seq_ctx"].cu_seq_lens_q.detach().cpu(), + "num_padding": data["seq_ctx"].num_padding, + "shifted_labels": data["shifted_labels"].detach().cpu().reshape(-1), + } + for data in data_batches + ] + captured_batch_index = 0 + apply_opd = worker_module.apply_opd_kl_to_advantages + + def capture_opd_result(loss_ctx, *, config): + nonlocal captured_batch_index + apply_opd(loss_ctx, config=config) + captured_batches[captured_batch_index]["old_log_probs"] = ( + loss_ctx.loss_kwargs.old_logprobs.detach().cpu().reshape(-1) + ) + captured_batches[captured_batch_index]["advantages"] = ( + loss_ctx.loss_kwargs.advantages.detach().cpu().reshape(-1) + ) + captured_batch_index += 1 - async def generate_sample(state, **kwargs): - return self._complete(state) + with mock_patch.object(worker_module, "apply_opd_kl_to_advantages", capture_opd_result): + worker_log_item = TrainingWorker.fit(self, data_batches, rollout_idx) - loop.generate_sample = generate_sample - result = await loop.collect_rollout_group([_state(1)], is_valid_sample_func=filter_func) + torch.save(captured_batches, self.trainer_capture_dir / f"rank_{self.rank}.pt") + return worker_log_item - self.assertEqual(order, ["filter", "teacher"]) - self.assertEqual(result[0].status, Status.COMPLETED) - self.assertIsNotNone(result[0].teacher_logprobs) - self.assertIn("teacher_score_time_s", result[0].extra_fields) + def build_capturing_training_workers(self, placement_group): + from xtuner.v1.rl.utils import AutoAcceleratorWorkers - async def test_origin_data_source_routes_groups_to_different_teachers(self): - math_endpoint = self._endpoint() - code_endpoint = self._endpoint() - loop = self._loop( - _opd_config( - {"math_teacher": math_endpoint.url, "code_teacher": code_endpoint.url}, - {"math": "math_teacher", "code": "code_teacher"}, - ) + capturing_worker_cls = ray.remote( + runtime_env={ + "env_vars": { + "RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES": "1", + "RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES": "1", + "HCCL_NPU_SOCKET_PORT_RANGE": "auto", + } + } + )(CapturingTrainingWorker) + train_workers, _ = AutoAcceleratorWorkers.from_placement_group( + capturing_worker_cls, + self, + placement_group, ) - - async def generate_sample(state, **kwargs): - return self._complete(state) - - loop.generate_sample = generate_sample - await loop.collect_rollout_group([_state(1, data_source="math")]) - await loop.collect_rollout_group([_state(2, data_source="code")]) - - self.assertEqual(len(math_endpoint.state.requests), 1) - self.assertEqual(len(code_endpoint.state.requests), 1) - - async def test_teacher_failure_marks_group_failed(self): - endpoint = self._endpoint([_HTTPResponse({"error": "down"}, status=500)]) - loop = self._loop( - _opd_config( - {"teacher": endpoint.url}, - {"math": "teacher"}, - max_retry_per_sample=0, + ray.wait([worker.ready.remote() for worker in train_workers]) + return CapturingTrainingController(workers=train_workers) + + trainer_environment = { + "WORK_DIR": str(run_dir / "trainer"), + "MODEL_PATH": str(student_model_path), + "DATA_PATH": os.getenv("XTUNER_OPD_DATA_PATH", str(baseline_path)), + "WORLD_SIZE": "1", + "ONLY_CALC_MISMATCH_RATIO": "1", + "XTUNER_DETERMINISTIC": "true", + "XTUNER_USE_FA3": os.getenv("XTUNER_OPD_USE_FA3", "1"), + "TOKENIZERS_PARALLELISM": "false", + "PYTHONUNBUFFERED": "1", + } + + with patch.dict(os.environ, trainer_environment, clear=False): + try: + ray.init( + num_cpus=12 * num_workers + 10, + num_gpus=num_workers, + include_dashboard=False, + _temp_dir="/dev/shm/xtuner-opd-old-logprobs", ) - ) + cfg = Config.fromfile(TRAINER_CONFIG_PATH) + cfg.trainer.resources.num_workers = num_workers + cfg.trainer.total_train_steps = 1 + cfg.trainer.train_batch_size = len(baseline_samples) + cfg.trainer.train_worker_cfg.optimizer_steps = 1 + cfg.trainer.train_worker_cfg.pack_max_length = pack_max_length + cfg.trainer.debug_train = True + cfg.trainer.debug_rollout_dir = debug_rollout_dir + with ( + patch("xtuner.v1.train.rl_trainer.XTUNER_DETERMINISTIC", True), + patch( + "xtuner.v1.train.rl_trainer.RLColocateTrainer", + CapturingRLColocateTrainer, + ), + patch.object(WorkerConfig, "build", build_capturing_training_workers), + ): + trainer = cfg.trainer.build() + trainer.fit() + finally: + if ray.is_initialized(): + ray.shutdown() - async def generate_sample(state, **kwargs): - return self._complete(state) - - loop.generate_sample = generate_sample - result = await loop.collect_rollout_group([_state(1)]) - - self.assertEqual(get_group_status(result), Status.FAILED) - self.assertIsNone(result[0].teacher_logprobs) - self.assertIn("teacher_score_time_s", result[0].extra_fields) - - -class TestOPDRolloutData(unittest.IsolatedAsyncioTestCase): - async def test_rollout_state_round_trip_replay_and_reset_preserve_contract(self): - state = _state(1) - state.teacher_tokens = [201, 2] - state.teacher_logprobs = [-0.2, -0.3] - - restored = RolloutState.model_validate(state.model_dump()) - self.assertEqual(restored.teacher_tokens, state.teacher_tokens) - self.assertEqual(restored.teacher_logprobs, state.teacher_logprobs) - - replay_buffer = AsyncReplayBufferConfig().build() - await replay_buffer.put([restored], "math") - replayed = (await replay_buffer.get(1, "math", Status.COMPLETED))[0][0] - self.assertEqual(replayed.teacher_tokens, [201, 2]) - self.assertEqual(replayed.teacher_logprobs, [-0.2, -0.3]) - - reset_rollout_response(replayed) - self.assertIsNone(replayed.teacher_tokens) - self.assertIsNone(replayed.teacher_logprobs) - - -class TestPGOPDTokenAdvantages(unittest.TestCase): - @staticmethod - def _scored_state( - rollout_id: int, - *, - behavior_logprobs: list[float], - teacher_logprobs: list[float], - reward: float | None = None, - response_mask: list[int] | None = None, - ) -> RolloutState: - response_ids = list(range(100, 100 + len(behavior_logprobs))) - state = _state(rollout_id, response_ids=response_ids) - state.logprobs = behavior_logprobs - state.teacher_tokens = response_ids - state.teacher_logprobs = teacher_logprobs - state.reward = None if reward is None else {"score": reward} - state.response_mask = response_mask - return state - - def test_pure_opd_uses_token_delta_and_response_mask_without_reward(self): - state = self._scored_state( - 1, - behavior_logprobs=[-1.0, -2.0, -3.0], - teacher_logprobs=[-0.5, -2.5, -3.0], - response_mask=[1, 0, 1], - ) + return _load_trainer_samples(capture_dir) - advantages = compute_pg_opd_token_advantages( - [state], - config=_algorithm_config(task_adv_weight=0.0, opd_adv_weight=2.0), - task_adv_estimator=None, - ) - torch.testing.assert_close(advantages[0], torch.tensor([1.0, 0.0, 0.0])) - self.assertFalse(advantages[0].requires_grad) - - state.response_mask = [] - unmasked = compute_pg_opd_token_advantages( - [state], - config=_algorithm_config(), - task_adv_estimator=None, - ) - torch.testing.assert_close(unmasked[0], torch.tensor([0.5, -0.5, 0.0])) - - def test_mixed_opd_combines_per_sample_task_advantage(self): - first = self._scored_state( - 1, - behavior_logprobs=[-1.0], - teacher_logprobs=[-0.75], - reward=10.0, - ) - second = self._scored_state( - 2, - behavior_logprobs=[-1.0], - teacher_logprobs=[-1.25], - reward=10.0, - ) - third = self._scored_state( - 3, - behavior_logprobs=[-1.0], - teacher_logprobs=[-0.5], - reward=-1.0, +@unittest.skipUnless( + BASELINE_PATH and STUDENT_MODEL_PATH, + "XTUNER_OPD_BASELINE and XTUNER_OPD_STUDENT_MODEL are required", +) +class TestPGOPDOldLogprobs(unittest.TestCase): + def test_trainer_old_logprobs_and_advantage_error_propagation(self) -> None: + baseline_path = Path(str(BASELINE_PATH)).expanduser().resolve() + student_model_path = Path(str(STUDENT_MODEL_PATH)).expanduser().resolve() + baseline = torch.load(baseline_path, map_location="cpu", weights_only=True) + baseline_samples = baseline["samples"] + + work_root = Path( + os.getenv( + "XTUNER_OPD_OLD_LOGPROB_WORK_DIR", + str(REPO_ROOT / "work_dirs/test_pg_opd_old_logprobs"), + ) + ).expanduser() + run_dir = work_root / f"run_{time.strftime('%Y%m%d_%H%M%S')}_{os.getpid()}" + debug_rollout_dir = run_dir / "debug_rollout" + capture_dir = run_dir / "trainer_capture" + debug_rollout_dir.mkdir(parents=True) + capture_dir.mkdir() + + train_batch = _build_debug_rollout_batch(baseline_samples) + torch.save(train_batch, debug_rollout_dir / "debug_rollout_1.pt") + trainer_samples_by_rollout_id = _run_trainer_once( + baseline_samples, + baseline_path=baseline_path, + student_model_path=student_model_path, + debug_rollout_dir=debug_rollout_dir, + capture_dir=capture_dir, + run_dir=run_dir, ) - estimator = _FixedAdvantageEstimator([2.0, 2.0, -2.0]) - advantages = compute_pg_opd_token_advantages( - [first, second, third], - config=_algorithm_config(task_adv_weight=0.5, opd_adv_weight=2.0), - task_adv_estimator=estimator, - ) + result_path = run_dir / "result.pt" + result_samples = [] + + for baseline_sample in baseline_samples: + rollout_id = int(baseline_sample["sample_index"]) + trainer_sample = trainer_samples_by_rollout_id[rollout_id] + loss_mask = torch.as_tensor(baseline_sample["loss_mask"], dtype=torch.bool) + response_token_ids = torch.as_tensor(baseline_sample["response_token_ids"], dtype=torch.long)[loss_mask] + torch.testing.assert_close( + trainer_sample["response_token_ids"], + response_token_ids, + rtol=0, + atol=0, + ) + baseline_old_log_probs = torch.as_tensor( + baseline_sample["old_log_probs"], + dtype=torch.float32, + )[loss_mask] + trainer_old_log_probs = trainer_sample["old_log_probs"] + baseline_advantages = torch.as_tensor( + baseline_sample["advantages"], + dtype=torch.float32, + )[loss_mask] + trainer_advantages = trainer_sample["advantages"] + + old_logprob_error = trainer_old_log_probs - baseline_old_log_probs + advantage_error = trainer_advantages - baseline_advantages + propagation_residual = advantage_error + old_logprob_error + sample_num_tokens = old_logprob_error.numel() + sample_summary = { + "num_tokens": sample_num_tokens, + "old_logprobs_mean_abs_error": old_logprob_error.abs().mean().item(), + "old_logprobs_max_abs_error": old_logprob_error.abs().max().item(), + "advantages_mean_abs_error": advantage_error.abs().mean().item(), + "advantages_max_abs_error": advantage_error.abs().max().item(), + "propagation_mean_abs_error": propagation_residual.abs().mean().item(), + "propagation_max_abs_error": propagation_residual.abs().max().item(), + } + result_samples.append( + { + "sample_index": rollout_id, + "response_token_ids": response_token_ids, + "baseline_old_log_probs": baseline_old_log_probs, + "trainer_old_log_probs": trainer_old_log_probs, + "old_logprob_error": old_logprob_error, + "baseline_advantages": baseline_advantages, + "trainer_advantages": trainer_advantages, + "advantage_error": advantage_error, + "propagation_residual": propagation_residual, + "summary": sample_summary, + } + ) - torch.testing.assert_close(torch.cat(advantages), torch.tensor([1.5, 0.5, 0.0])) - self.assertEqual(estimator.calls[0][0].tolist(), [10.0, 10.0, -1.0]) - self.assertEqual(estimator.calls[0][1], [first, second, third]) + with self.subTest(sample_index=rollout_id): + torch.testing.assert_close( + advantage_error, + -old_logprob_error, + rtol=1e-5, + atol=1e-5, + msg=f"Result: {result_path}", + ) - def test_mixed_opd_requires_reward(self): - state = self._scored_state( - 1, - behavior_logprobs=[-1.0], - teacher_logprobs=[-0.5], + global_result = { + "response_token_ids": torch.cat([sample["response_token_ids"] for sample in result_samples]), + "baseline_old_log_probs": torch.cat([sample["baseline_old_log_probs"] for sample in result_samples]), + "trainer_old_log_probs": torch.cat([sample["trainer_old_log_probs"] for sample in result_samples]), + "old_logprob_error": torch.cat([sample["old_logprob_error"] for sample in result_samples]), + "baseline_advantages": torch.cat([sample["baseline_advantages"] for sample in result_samples]), + "trainer_advantages": torch.cat([sample["trainer_advantages"] for sample in result_samples]), + "advantage_error": torch.cat([sample["advantage_error"] for sample in result_samples]), + "propagation_residual": torch.cat([sample["propagation_residual"] for sample in result_samples]), + } + summary = { + "num_samples": len(result_samples), + "num_tokens": global_result["old_logprob_error"].numel(), + "old_logprobs_mean_abs_error": global_result["old_logprob_error"].abs().mean().item(), + "old_logprobs_max_abs_error": global_result["old_logprob_error"].abs().max().item(), + "advantages_mean_abs_error": global_result["advantage_error"].abs().mean().item(), + "advantages_max_abs_error": global_result["advantage_error"].abs().max().item(), + "propagation_mean_abs_error": global_result["propagation_residual"].abs().mean().item(), + "propagation_max_abs_error": global_result["propagation_residual"].abs().max().item(), + } + torch.save( + { + "summary": summary, + "global": global_result, + "samples": result_samples, + }, + result_path, ) - config = _algorithm_config(task_adv_weight=1.0) - with self.assertRaisesRegex(ValueError, "Reward score is required"): - compute_pg_opd_token_advantages( - [state], - config=config, - task_adv_estimator=_FixedAdvantageEstimator([1.0]), + with self.subTest(scope="global"): + torch.testing.assert_close( + global_result["advantage_error"], + -global_result["old_logprob_error"], + rtol=1e-5, + atol=1e-5, + msg=f"Result: {result_path}", ) - def test_token_advantages_survive_packing_and_define_denominator(self): - state = self._scored_state( - 1, - behavior_logprobs=[-1.0, -2.0, -3.0], - teacher_logprobs=[-0.5, -2.5, -2.0], - response_mask=[1, 0, 1], + print( + f"old_logprobs: mean_abs={summary['old_logprobs_mean_abs_error']}, " + f"max_abs={summary['old_logprobs_max_abs_error']}\n" + f"advantages: mean_abs={summary['advantages_mean_abs_error']}, " + f"max_abs={summary['advantages_max_abs_error']}\n" + f"propagation: mean_abs={summary['propagation_mean_abs_error']}, " + f"max_abs={summary['propagation_max_abs_error']}\n" + f"result: {result_path}" ) - response_advantages = compute_pg_opd_token_advantages( - [state], - config=_algorithm_config(), - task_adv_estimator=None, - )[0] - data_batches = [ - { - "seq_ctx": SequenceContext.from_input_ids((torch.tensor([[10, 11, 20, 21]]),), device="cpu"), - "shifted_labels": torch.tensor([[-100, 100, -100, 102]]), - "advantage": [0.0] + response_advantages.tolist(), - "rollout_logprobs": torch.tensor([[0.0, -1.0, -2.0, -3.0]]), - } - ] - - controller = TrainingController.__new__(TrainingController) - packed = controller._packing(data_batches, pack_max_length=6, language_cfg=None) - self.assertEqual(packed[0]["shifted_labels"].shape, packed[0]["advantages"].shape) - self.assertEqual(packed[0]["advantages"].tolist(), [[0.0, 0.5, 0.0, 1.0, -100.0, -100.0]]) - loss_config = GRPOLossConfig( - policy_loss_cfg={"loss_type": "vanilla", "cliprange_low": 0.2, "cliprange_high": 0.2} - ) - loss_context = loss_config.build( - { - "shifted_labels": packed[0]["shifted_labels"], - "advantages": packed[0]["advantages"], - "old_logprobs": torch.zeros_like(packed[0]["advantages"]), - } - ) - assert isinstance(loss_context, GRPOLossContext) - GRPOLossContext.build_batches([loss_context]) - torch.testing.assert_close( - loss_context.loss_kwargs.policy_loss_weight.cpu(), - torch.tensor([[0.0, 0.5, 0.0, 0.5, 0.0, 0.0]]), +@unittest.skipUnless( + BASELINE_PATH and TEACHER_MODEL_PATH, + "XTUNER_OPD_BASELINE and XTUNER_OPD_TEACHER_MODEL are required", +) +class TestTeacherLogprobClient(unittest.IsolatedAsyncioTestCase): + @classmethod + def setUpClass(cls) -> None: + super().setUpClass() + baseline = torch.load(BASELINE_PATH, map_location="cpu", weights_only=False) + cls.samples = baseline["samples"] + port = find_free_ports()[0] + cls.teacher_endpoint = f"http://127.0.0.1:{port}" + teacher_env = os.environ.copy() + teacher_env["TEACHER_HOST"] = "127.0.0.1" + teacher_env["TEACHER_PORT"] = str(port) + cls.teacher_process = subprocess.Popen( + ["bash", str(TEACHER_START_SCRIPT), str(TEACHER_MODEL_PATH)], + cwd=REPO_ROOT, + env=teacher_env, + start_new_session=True, ) + cls.addClassCleanup(cls._stop_teacher) + _wait_for_teacher(cls.teacher_process, cls.teacher_endpoint) + + @classmethod + def _stop_teacher(cls) -> None: + if cls.teacher_process.poll() is not None: + return + os.killpg(cls.teacher_process.pid, signal.SIGTERM) + try: + cls.teacher_process.wait(timeout=30) + except subprocess.TimeoutExpired: + os.killpg(cls.teacher_process.pid, signal.SIGKILL) + cls.teacher_process.wait() + + async def test_compute_logprobs_matches_baseline(self) -> None: + client = TeacherLogprobClient(OPDTeacherConfig(name="teacher", endpoint=self.teacher_endpoint)) + self.addAsyncCleanup(client._client.aclose) + + for sample in self.samples: + with self.subTest(sample_index=sample["sample_index"]): + prompt_ids = torch.as_tensor(sample["prompt_token_ids"], dtype=torch.long).tolist() + response_ids = torch.as_tensor(sample["response_token_ids"], dtype=torch.long).tolist() + expected_logprobs = torch.as_tensor(sample["teacher_log_probs"], dtype=torch.float32) + state = RolloutState( + rollout_id=int(sample["sample_index"]), + group_id=int(sample["group_index"]), + message=[], + prompt_ids=prompt_ids, + tokens=prompt_ids, + response="", + response_ids=response_ids, + status=Status.COMPLETED, + ) - def test_teacher_signal_changes_student_gradient(self): - def gradient(teacher_logprob: float) -> torch.Tensor: - state = self._scored_state( - 1, - behavior_logprobs=[-1.0], - teacher_logprobs=[teacher_logprob], - ) - advantage = compute_pg_opd_token_advantages( - [state], - config=_algorithm_config(), - task_adv_estimator=None, - )[0].unsqueeze(0) - current_logprob = torch.tensor([[-1.0]], requires_grad=True) - old_logprob = torch.tensor([[-1.0]]) - loss_config = GRPOLossConfig( - policy_loss_cfg={"loss_type": "vanilla", "cliprange_low": 0.2, "cliprange_high": 0.2} - ) - context = loss_config.build( - { - "shifted_labels": torch.tensor([[100]]), - "advantages": advantage, - "old_logprobs": old_logprob, - } - ) - assert isinstance(context, GRPOLossContext) - GRPOLossContext.build_batches([context]) - device = context.loss_kwargs.advantages.device - current_logprob = current_logprob.to(device).detach().requires_grad_() - loss = context.policy_loss_fn( - current_logprob, - context.loss_kwargs.old_logprobs, - context.loss_kwargs.advantages, - context.loss_kwargs.policy_loss_weight, - loss_config.policy_loss_cfg, - ) - loss.backward() - return current_logprob.grad.detach().cpu().clone() - - positive_signal_gradient = gradient(-0.5) - stronger_signal_gradient = gradient(0.0) - - self.assertFalse(torch.equal(positive_signal_gradient, stronger_signal_gradient)) - - def test_pure_and_mixed_opd_support_rollout_is_off_and_on(self): - for task_adv_weight in (0.0, 1.0): - for enable_is in (False, True): - with self.subTest(task_adv_weight=task_adv_weight, enable_is=enable_is): - state = self._scored_state( - 1, - behavior_logprobs=[-1.0, -1.0], - teacher_logprobs=[0.0, 0.0], - reward=1.0 if task_adv_weight else None, + result = await client.compute_logprobs(state) + + self.assertEqual(result.status, Status.COMPLETED, result.error_msg) + self.assertEqual(result.teacher_tokens, response_ids) + actual_logprobs = torch.tensor(result.teacher_logprobs, dtype=torch.float32) + try: + torch.testing.assert_close( + actual_logprobs, + expected_logprobs, + rtol=1e-5, + atol=1e-5, ) - estimator = _FixedAdvantageEstimator([1.0]) if task_adv_weight else None - advantages = compute_pg_opd_token_advantages( - [state], - config=_algorithm_config(task_adv_weight=task_adv_weight), - task_adv_estimator=estimator, - )[0].unsqueeze(0) - rollout_is = ( - RolloutImportanceSampling( - rollout_is_mode="mask", - rollout_is_threshold=(2.0, 0.5), - rollout_is_mask_threshold=(2.0, 0.5), + except AssertionError as error: + mismatch_indices = torch.nonzero( + ~torch.isclose(actual_logprobs, expected_logprobs, rtol=1e-5, atol=1e-5) + ).flatten() + mismatch_values = "\n".join( + ( + f"index={index}: " + f"actual={actual_logprobs[index].item()!r}, " + f"expected={expected_logprobs[index].item()!r}, " + f"abs_diff={abs(actual_logprobs[index] - expected_logprobs[index]).item()!r}" ) - if enable_is - else RolloutImportanceSampling() - ) - loss_config = GRPOLossConfig( - policy_loss_cfg={"loss_type": "vanilla", "cliprange_low": 0.2, "cliprange_high": 0.2}, - rollout_is=rollout_is, + for index in mismatch_indices.tolist() ) - context = loss_config.build( - { - "shifted_labels": torch.tensor([[100, 101]]), - "advantages": advantages, - "rollout_logprobs": torch.tensor([[-1.0, -1.0]]), - "old_logprobs": torch.tensor([[-1.0, -1.0 + math.log(10.0)]]), - } - ) - assert isinstance(context, GRPOLossContext) - context.compute_rollout_is(None, torch.tensor([2])) # type: ignore[arg-type] - GRPOLossContext.build_batches([context]) - - device = context.loss_kwargs.advantages.device - current_logprobs = torch.tensor([[-1.0, -1.0]], device=device, requires_grad=True) - loss = context.policy_loss_fn( - current_logprobs, - context.loss_kwargs.old_logprobs, - context.loss_kwargs.advantages, - context.loss_kwargs.policy_loss_weight, - loss_config.policy_loss_cfg, - ) - loss.backward() - - if enable_is: - self.assertIsNotNone(context.loss_kwargs.is_weights) - self.assertEqual(context.loss_kwargs.shifted_labels.tolist(), [[100, -100]]) - self.assertEqual(current_logprobs.grad[0, 1].item(), 0.0) - else: - self.assertIsNone(context.loss_kwargs.is_weights) - self.assertEqual(context.loss_kwargs.shifted_labels.tolist(), [[100, 101]]) - self.assertNotEqual(current_logprobs.grad[0, 1].item(), 0.0) - - -class TestOPDConfig(unittest.TestCase): - def test_rejects_duplicate_or_unknown_teacher_names(self): - teacher = OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:1") - with self.assertRaisesRegex(ValueError, "must be unique"): - OPDConfig( - teachers=[teacher, teacher], - data_source_teacher_map={"math": "teacher"}, - ) - with self.assertRaisesRegex(ValueError, "unknown teachers"): - OPDConfig( - teachers=[teacher], - data_source_teacher_map={"math": "missing"}, - ) + raise AssertionError(f"{error}\n\nMismatched logprobs:\n{mismatch_values}") from None if __name__ == "__main__": diff --git a/tests/rl/test_prepare_train_data.py b/tests/rl/test_prepare_train_data.py index 81a364c889..efb32d61c5 100644 --- a/tests/rl/test_prepare_train_data.py +++ b/tests/rl/test_prepare_train_data.py @@ -7,9 +7,6 @@ - VLM 样本使用 train_prompt_ids,并保留 multimodal 训练字段。 - 无效 rollout group 会被跳过。 - 缺失 reward、logprob/mask 长度不一致、pack_max_length 过小时 fail fast。 - -注意:当前训练 contract 中 data_dict["advantage"] 比 shifted_labels 多 1 个元素; -metric 统计使用 actual_advantages[:-1],测试会显式固定这个行为。 """ import unittest @@ -19,7 +16,6 @@ import torch from xtuner.v1.data_proto.rl_data import RolloutState, Status -from xtuner.v1.rl.on_policy_distillation import OPDConfig, OPDTeacherConfig from xtuner.v1.train.rl_trainer import BaseRLTrainer @@ -34,23 +30,14 @@ def compute(self, rewards_tensor, group): class TestPrepareTrainData(unittest.TestCase): - def _build_trainer(self, advantages: list[float], opd_config: OPDConfig | None = None): + def _build_trainer(self, advantages: list[float]): trainer = BaseRLTrainer.__new__(BaseRLTrainer) trainer._advantage_estimator = _FakeAdvantageEstimator(advantages) - trainer._opd_config = opd_config + trainer._opd_config = None trainer.tokenizer = MagicMock(return_value={"input_ids": torch.tensor([[999]])}) trainer.logger = MagicMock() return trainer - @staticmethod - def _opd_config(*, task_adv_weight: float = 0.0, opd_adv_weight: float = 1.0) -> OPDConfig: - return OPDConfig( - task_adv_weight=task_adv_weight, - opd_adv_weight=opd_adv_weight, - teachers=[OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:1")], - data_source_teacher_map={"math": "teacher"}, - ) - def _state( self, *, @@ -113,8 +100,8 @@ def test_text_path_builds_shifted_training_tensors(self): batch["rollout_logprobs"], torch.tensor([[0.0, 0.0, 0.1, 0.2, 0.3]], dtype=torch.float32), ) - self.assertEqual(batch["advantage"], [1.5, 1.5, 1.5, 1.5, 0.0, 1.5]) - self.assertEqual(len(batch["advantage"]), batch["shifted_labels"].numel() + 1) + self.assertEqual(batch["advantage"], [0.0, 0.0, 1.5, 0.0, 1.5]) + self.assertEqual(len(batch["advantage"]), batch["shifted_labels"].numel()) self.assertIs(batch["seq_ctx"].rollout_routed_experts, routed_experts) self.assertEqual(info["training_samples"], 1) self.assertEqual(info["training_tokens"], 5) @@ -131,8 +118,8 @@ def test_multi_sample_group_uses_each_sample_reward_and_advantage(self): data_batches, info = self._prepare(trainer, [[first, second]]) self.assertEqual(len(data_batches), 2) - self.assertEqual(data_batches[0]["advantage"], [1.5, 1.5, 1.5, 1.5, 1.5]) - self.assertEqual(data_batches[1]["advantage"], [-2.0, -2.0, -2.0, -2.0, -2.0]) + self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 1.5, 1.5]) + self.assertEqual(data_batches[1]["advantage"], [0.0, 0.0, -2.0, -2.0]) self.assertEqual(info["batch_size"], 2) self.assertEqual(info["rewards/min"], -1.0) self.assertEqual(info["rewards/max"], 3.0) @@ -141,68 +128,6 @@ def test_multi_sample_group_uses_each_sample_reward_and_advantage(self): self.assertEqual(info["advantages/max"], 1.5) self.assertEqual(trainer._advantage_estimator.calls[0][0].tolist(), [3.0, -1.0]) - def test_pure_opd_builds_aligned_token_advantages_without_reward(self): - trainer = self._build_trainer([99.0], self._opd_config()) - state = self._state( - prompt_ids=[10, 11], - response_ids=[20, 21, 2], - logprobs=[-1.0, -2.0, -3.0], - response_mask=[1, 0, 1], - ) - state.reward = None - state.teacher_tokens = [20, 21, 2] - state.teacher_logprobs = [-0.5, -2.5, -2.0] - - data_batches, info = self._prepare(trainer, [[state]]) - - self.assertEqual(data_batches[0]["seq_ctx"].input_ids.tolist(), [[10, 11, 20, 21]]) - self.assertEqual(data_batches[0]["shifted_labels"].tolist(), [[-100, 20, -100, 2]]) - self.assertEqual(data_batches[0]["advantage"], [0.0, 0.5, 0.0, 1.0]) - self.assertEqual(len(data_batches[0]["advantage"]), data_batches[0]["shifted_labels"].numel()) - self.assertEqual(trainer._advantage_estimator.calls, []) - self.assertEqual(info["training_samples"], 1) - self.assertEqual(info["rewards/mean"], 0.0) - - def test_mixed_opd_combines_cluster_task_and_token_advantages(self): - trainer = self._build_trainer( - [2.0, -1.0], - self._opd_config(task_adv_weight=0.5, opd_adv_weight=2.0), - ) - first = self._state( - uid=1, - response_ids=[20, 2], - logprobs=[-1.0, -1.0], - reward={"score": 3.0}, - ) - first.teacher_tokens = [20, 2] - first.teacher_logprobs = [-0.5, -1.0] - second = self._state( - uid=2, - response_ids=[30, 2], - logprobs=[-1.0, -1.0], - reward={"score": -1.0}, - ) - second.teacher_tokens = [30, 2] - second.teacher_logprobs = [-1.5, -0.5] - - data_batches, info = self._prepare(trainer, [[first, second]]) - - self.assertEqual(data_batches[0]["advantage"], [0.0, 0.0, 2.0, 1.0]) - self.assertEqual(data_batches[1]["advantage"], [0.0, 0.0, -1.5, 0.5]) - self.assertEqual(trainer._advantage_estimator.calls[0][0].tolist(), [3.0, -1.0]) - self.assertEqual(info["rewards/min"], -1.0) - self.assertEqual(info["rewards/max"], 3.0) - - def test_mixed_opd_missing_reward_fails_fast(self): - trainer = self._build_trainer([1.0], self._opd_config(task_adv_weight=1.0)) - state = self._state(response_ids=[20], logprobs=[-1.0]) - state.reward = None - state.teacher_tokens = [20] - state.teacher_logprobs = [-0.5] - - with self.assertRaisesRegex(ValueError, "Reward score is required"): - self._prepare(trainer, [[state]]) - 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]) diff --git a/xtuner/v1/rl/loss/base_loss.py b/xtuner/v1/rl/loss/base_loss.py index 108af41ed8..18fc2d7175 100644 --- a/xtuner/v1/rl/loss/base_loss.py +++ b/xtuner/v1/rl/loss/base_loss.py @@ -106,6 +106,7 @@ def build( - shifted_labels (torch.Tensor): The shifted labels - advantages (torch.Tensor): Advantage estimates - rollout_logprobs (torch.Tensor | None): Rollout log probabilities + - teacher_logprobs (torch.Tensor | None): Teacher log probabilities for OPD - old_logprobs (torch.Tensor | None): Old policy log probabilities (optional, can be set later) - rollout_is_weights (torch.Tensor | None): Importance sampling weights - ref_logprobs (torch.Tensor | None): Reference model log probabilities @@ -122,6 +123,7 @@ def build( shifted_labels = data["shifted_labels"] advantages = data["advantages"] rollout_logprobs = data.get("rollout_logprobs", None) + teacher_logprobs = data.get("teacher_logprobs", None) old_logprobs = data.get("old_logprobs", None) rollout_is_weights = data.get("rollout_is_weights", None) ref_logprobs = data.get("ref_logprobs", None) @@ -132,6 +134,7 @@ def build( old_logprobs=old_logprobs, advantages=advantages, rollout_logprobs=rollout_logprobs, + teacher_logprobs=teacher_logprobs, is_weights=rollout_is_weights, ref_logprobs=ref_logprobs, ).to(DEVICE) @@ -153,10 +156,12 @@ class BaseRLLossKwargs(CELossKwargs): ref_logprobs (torch.Tensor | None): Reference log probabilities for KL penalty, if used. kl_loss_weight (torch.Tensor | None): Weights for each token in the KL loss computation, if used. rollout_logprobs (torch.Tensor | None): Rollout log probabilities from inference engine, used for importance sampling. + teacher_logprobs (torch.Tensor | None): Teacher log probabilities used to apply the OPD KL penalty. is_weights (torch.Tensor | None): Importance sampling weights. If None, importance sampling is not used. """ rollout_logprobs: torch.Tensor | None = None + teacher_logprobs: torch.Tensor | None = None advantages: torch.Tensor old_logprobs: torch.Tensor | None = None policy_loss_weight: torch.Tensor | None = None @@ -172,6 +177,8 @@ def sp_split(self, sp_mesh: DeviceMesh) -> Self: self.advantages = sp_split(self.advantages, sp_mesh=sp_mesh, split_dim=1, padding_value=0.0) if self.rollout_logprobs is not None: self.rollout_logprobs = sp_split(self.rollout_logprobs, sp_mesh=sp_mesh, split_dim=1, padding_value=0.0) + if self.teacher_logprobs is not None: + self.teacher_logprobs = sp_split(self.teacher_logprobs, sp_mesh=sp_mesh, split_dim=1, padding_value=0.0) if self.is_weights is not None: self.is_weights = sp_split(self.is_weights, sp_mesh=sp_mesh, split_dim=1, padding_value=1.0) # 1. 这里不用对old_logprobs和ref_logprobs进行sp_split,因为他是模型 fwd 生成的, @@ -192,6 +199,8 @@ def to(self, device: torch.device | str) -> Self: self.ref_logprobs = self.ref_logprobs.to(device) if self.rollout_logprobs is not None: self.rollout_logprobs = self.rollout_logprobs.to(device) + if self.teacher_logprobs is not None: + self.teacher_logprobs = self.teacher_logprobs.to(device) if self.is_weights is not None: self.is_weights = self.is_weights.to(device) if self.global_grad_tokens is not None: diff --git a/xtuner/v1/rl/on_policy_distillation.py b/xtuner/v1/rl/on_policy_distillation.py index 739495ea32..62d8961f95 100644 --- a/xtuner/v1/rl/on_policy_distillation.py +++ b/xtuner/v1/rl/on_policy_distillation.py @@ -10,7 +10,7 @@ from pydantic import BaseModel, ConfigDict, Field, model_validator from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status -from xtuner.v1.rl.advantage import AdvantageEstimator +from xtuner.v1.rl.loss.base_loss import BaseRLLossContext class OPDTeacherConfig(BaseModel): @@ -88,11 +88,11 @@ async def compute_logprobs(self, state: RolloutState) -> RolloutState: "input_ids": prompt_ids + response_ids, "sampling_params": { "max_new_tokens": 0, - "temperature": 1.0, + "temperature": 0, "skip_special_tokens": False, }, "return_logprob": True, - "logprob_start_len": max(len(prompt_ids) - 1, 0), + "logprob_start_len": 0, "top_logprobs_num": 0, "stream": False, } @@ -124,12 +124,13 @@ def _parse_response( ) -> tuple[list[int], list[float]]: try: raw_logprobs = response.json()["meta_info"]["input_token_logprobs"] - teacher_tokens = [item[1] for item in raw_logprobs[1:]] - teacher_logprobs = [float(item[0]) for item in raw_logprobs[1:]] + response_logprobs = raw_logprobs[-len(response_ids) :] + teacher_tokens = [item[1] for item in response_logprobs] + teacher_logprobs = [float(item[0]) for item in response_logprobs] except (KeyError, TypeError, IndexError) as exc: raise ValueError("Invalid teacher response") from exc - if len(raw_logprobs) != len(response_ids) + 1: + if len(teacher_logprobs) != len(response_ids): raise ValueError("Teacher logprob length mismatch") if teacher_tokens != response_ids: raise ValueError("Teacher token ids mismatch") @@ -149,52 +150,22 @@ def route_teacher_client( return teacher_clients[teacher_name] -def compute_pg_opd_token_advantages( - group: list[RolloutState], +def apply_opd_kl_to_advantages( + loss_ctx: BaseRLLossContext, *, config: OPDConfig, - task_adv_estimator: AdvantageEstimator | None, -) -> list[torch.Tensor]: - opd_advantages: list[torch.Tensor] = [] - response_masks: list[torch.Tensor] = [] - for state in group: - response_ids = cast(list[int], state.response_ids) - behavior_logprobs = cast(list[float], state.logprobs) - teacher_logprobs = cast(list[float], state.teacher_logprobs) - - behavior_logprobs_t = torch.tensor(behavior_logprobs, dtype=torch.float32) - teacher_logprobs_t = torch.tensor(teacher_logprobs, dtype=torch.float32) - response_mask = state.response_mask - if not response_mask: - response_mask_t = torch.ones(len(response_ids), dtype=torch.float32) - else: - response_mask_t = torch.tensor(response_mask, dtype=torch.float32) - opd_advantages.append(teacher_logprobs_t - behavior_logprobs_t) - response_masks.append(response_mask_t) - - task_advantages = [0.0] * len(group) - if config.task_adv_weight > 0: - rewards: list[float] = [] - for state in group: - if state.reward is None or "score" not in state.reward: - raise ValueError(f"Reward score is required for mixed PG-OPD rollout {state.rollout_id}") - rewards.append(float(state.reward["score"])) - - task_advantages = ( - cast(AdvantageEstimator, task_adv_estimator) - .compute( - torch.tensor(rewards, dtype=torch.float32), - group, - ) - .tolist() - ) - - return [ - (config.task_adv_weight * task_advantage + config.opd_adv_weight * opd_advantage) * response_mask - for task_advantage, opd_advantage, response_mask in zip( - task_advantages, - opd_advantages, - response_masks, - strict=True, - ) - ] +) -> torch.Tensor: + """Apply the OPD reverse-KL penalty and return its valid-token sum.""" + loss_kwargs = loss_ctx.loss_kwargs + old_logprobs = cast(torch.Tensor, loss_kwargs.old_logprobs) + teacher_logprobs = cast(torch.Tensor, loss_kwargs.teacher_logprobs) + response_mask = loss_kwargs.shifted_labels != loss_ctx.loss_cfg.ignore_idx + reverse_kl = old_logprobs - teacher_logprobs + reverse_kl_sum = (reverse_kl * response_mask).sum().detach() + loss_kwargs.advantages = torch.where( + response_mask, + loss_kwargs.advantages - config.opd_adv_weight * reverse_kl, + loss_kwargs.advantages, + ) + loss_kwargs.teacher_logprobs = None + return reverse_kl_sum diff --git a/xtuner/v1/rl/trainer/controller.py b/xtuner/v1/rl/trainer/controller.py index b87da88ced..e7e8b06e24 100644 --- a/xtuner/v1/rl/trainer/controller.py +++ b/xtuner/v1/rl/trainer/controller.py @@ -22,6 +22,7 @@ class ColateItem(TypedDict): shifted_labels: torch.Tensor advantage: float rollout_logprobs: torch.Tensor | None + teacher_logprobs: torch.Tensor | None class TrainingController: @@ -98,6 +99,10 @@ def _packing(self, data_batches, pack_max_length, language_cfg): if "rollout_logprobs" in data_batches[0] and data_batches[0]["rollout_logprobs"] is not None: rollout_logprobs_list = [data_batches[i]["rollout_logprobs"] for i in indices] + teacher_logprobs_list = None + if "teacher_logprobs" in data_batches[0] and data_batches[0]["teacher_logprobs"] is not None: + teacher_logprobs_list = [data_batches[i]["teacher_logprobs"] for i in indices] + if pad_len > 0: # Reduce the attn calculation time by using multiple short sequence packs pad_tokens = tuple( @@ -139,6 +144,14 @@ def _packing(self, data_batches, pack_max_length, language_cfg): device=data_batches[0]["shifted_labels"].device, ) rollout_logprobs_list.append(pad_rollout_logprobs) + if teacher_logprobs_list is not None: + pad_teacher_logprobs = torch.zeros( + 1, + pad_len, + dtype=data_batches[0]["teacher_logprobs"].dtype, + device=data_batches[0]["shifted_labels"].device, + ) + teacher_logprobs_list.append(pad_teacher_logprobs) seq_ctx = SequenceContext.cat(seq_ctx_list) shifted_labels = torch.cat(label_list, dim=1) # (1, max_len) @@ -149,12 +162,17 @@ def _packing(self, data_batches, pack_max_length, language_cfg): if rollout_logprobs_list is not None: rollout_logprobs = torch.cat(rollout_logprobs_list, dim=1) # (1, max_len) + teacher_logprobs = None + if teacher_logprobs_list is not None: + teacher_logprobs = torch.cat(teacher_logprobs_list, dim=1) # (1, max_len) + packed_data_batches.append( { "seq_ctx": seq_ctx, "shifted_labels": shifted_labels, "advantages": advantages, "rollout_logprobs": rollout_logprobs, + "teacher_logprobs": teacher_logprobs, } ) return packed_data_batches @@ -237,11 +255,17 @@ def fit(self, data_batches: list[ColateItem], pack_max_length: int, rollout_idx: pad_rollout_logprobs = torch.zeros( 1, pack_max_length, dtype=packed_data_batches[0]["rollout_logprobs"].dtype, device="cpu" ) + pad_teacher_logprobs = None + if "teacher_logprobs" in packed_data_batches[0] and packed_data_batches[0]["teacher_logprobs"] is not None: + pad_teacher_logprobs = torch.zeros( + 1, pack_max_length, dtype=packed_data_batches[0]["teacher_logprobs"].dtype, device="cpu" + ) pad_data = { "seq_ctx": pad_seq_ctx, "shifted_labels": pad_shifted_labels, "advantages": pad_advantages, "rollout_logprobs": pad_rollout_logprobs, + "teacher_logprobs": pad_teacher_logprobs, } pad_data_samples = [pad_data for _ in range(pad_num)] packed_data_batches = packed_data_batches + pad_data_samples diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index f7e0bc4871..1848866d1c 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -45,6 +45,7 @@ from xtuner.v1.model.utils.misc import ModelForwardExtraLogInfo from xtuner.v1.profiler import profiling_memory, profiling_time from xtuner.v1.rl.loss import BaseRLLossConfig, BaseRLLossContext, finalize_train_policy_metrics, kl_penalty +from xtuner.v1.rl.on_policy_distillation import OPDConfig, apply_opd_kl_to_advantages from xtuner.v1.rl.utils import SingleAcceleratorWorker from xtuner.v1.rl.weight_update import UpdateWeighter from xtuner.v1.train.trainer import LoadCheckpointConfig @@ -151,6 +152,7 @@ class WorkerConfig(BaseModel): profile_memory: bool = False free_rollout_routed_experts_in_worker: bool = True # 默认不需要用户配置 offload_rollout_routed_experts: bool = False + opd_config: OPDConfig | None = None # sft config sft_dataloader_cfg: DataloaderConfig | None = None @@ -184,6 +186,7 @@ class WorkerInputItem(TypedDict): shifted_labels: torch.LongTensor advantages: torch.Tensor rollout_logprobs: torch.Tensor | None + teacher_logprobs: torch.Tensor | None class WorkerTrainLogItem(TypedDict, total=False): @@ -196,6 +199,7 @@ class WorkerTrainLogItem(TypedDict, total=False): class WorkerLogItem(TypedDict): train_entropy: float + opd_reverse_kl: NotRequired[float] rollout_entropy: NotRequired[float] mismatch_metrics: NotRequired[dict[str, float]] rollout_is_metrics: NotRequired[dict[str, float]] @@ -586,6 +590,7 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo "shifted_labels": shifted_labels, "advantages": advantages, "rollout_logprobs": rollout_logprobs, + "teacher_logprobs": data.get("teacher_logprobs", None), }, sp_mesh=self.sp_mesh, ) @@ -625,8 +630,14 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo # compute old logprobs old_logprobs_list = self.compute_actor_logprobs(seq_ctx_list, shifted_labels_list) + rank_opd_reverse_kl_sum: torch.Tensor | None = None for old_logprobs, loss_ctx in zip(old_logprobs_list, loss_ctx_list): loss_ctx.loss_kwargs.old_logprobs = old_logprobs + if self.config.opd_config is not None: + reverse_kl_sum = apply_opd_kl_to_advantages(loss_ctx, config=self.config.opd_config) + rank_opd_reverse_kl_sum = ( + reverse_kl_sum if rank_opd_reverse_kl_sum is None else rank_opd_reverse_kl_sum + reverse_kl_sum + ) worker_log_item: WorkerLogItem = {"train_entropy": 0.0, "train_metrics": [], "sft_train_metrics": {}} logger_msg = f"Rollout {rollout_idx}: " @@ -651,6 +662,17 @@ def fit(self, data_batches: list[WorkerInputItem], rollout_idx: int) -> WorkerLo worker_log_item["rollout_entropy"] = avg_rollout_entropy.item() logger_msg += f", avg rollout entropy: {avg_rollout_entropy:.4f}" + if rank_opd_reverse_kl_sum is not None: + global_opd_reverse_kl_sum = rank_opd_reverse_kl_sum + dist.all_reduce(global_opd_reverse_kl_sum, op=dist.ReduceOp.SUM) + avg_opd_reverse_kl = ( + global_opd_reverse_kl_sum / global_grad_tokens + if global_grad_tokens > 0 + else global_opd_reverse_kl_sum.new_zeros(()) + ) + worker_log_item["opd_reverse_kl"] = avg_opd_reverse_kl.item() + logger_msg += f", OPD reverse KL: {avg_opd_reverse_kl:.4f}" + # compute rollout importance sampling metrics all_rollout_is_metrics = [] all_mismatch_metrics = [] diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index e650ad1ca3..6e7f003e3b 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -33,10 +33,7 @@ ) from xtuner.v1.rl.agent_loop_manager.produce_utils import default_should_continue_fn from xtuner.v1.rl.evaluator import EvaluatorConfig -from xtuner.v1.rl.on_policy_distillation import ( - OPDConfig, - compute_pg_opd_token_advantages, -) +from xtuner.v1.rl.on_policy_distillation import OPDConfig from xtuner.v1.rl.replay_buffer import ( AsyncReplayBufferConfig, SyncReplayBufferConfig, @@ -681,6 +678,7 @@ def _init_train_worker_config(self, cfg: BaseRLTrainerConfig, log_dir: Path) -> cfg.train_worker_cfg.free_rollout_routed_experts_in_worker = False cfg.train_worker_cfg.load_from = cfg.load_from cfg.train_worker_cfg.log_dir = log_dir + cfg.train_worker_cfg.opd_config = cfg.opd_config self._train_worker_cfg = cfg.train_worker_cfg def _init_rollout_config(self, cfg: BaseRLTrainerConfig, log_dir: Path) -> None: @@ -1041,6 +1039,7 @@ def _prepare_train_data( # When opd_config is not None, PG-OPD currently supports only single-turn rollouts. opd_config = self._opd_config + task_adv_weight = opd_config.task_adv_weight if opd_config is not None else 1.0 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 @@ -1058,19 +1057,8 @@ def _prepare_train_data( if isinstance(turns, int): tool_turns_list.append(turns) - opd_token_advantages: list[torch.Tensor] | None = None - if opd_config is not None: - opd_token_advantages = compute_pg_opd_token_advantages( - group, - config=opd_config, - task_adv_estimator=self._advantage_estimator, - ) - sample_advantages: list[float] = [] - - if opd_config.task_adv_weight > 0: - rewards = [float(cast(dict[str, Any], data.reward)["score"]) for data in group] - rewards_list.extend(rewards) - cluster_rewards_list.extend(rewards) + if task_adv_weight == 0: + sample_advantages = [0.0] * len(group) else: rewards = [] # Agentic rollouts may split one model session into multiple trainable segments. @@ -1103,6 +1091,7 @@ def _prepare_train_data( sample_advantages = [ cluster_advantages[cluster_index].item() for cluster_index in sample_cluster_indices ] + sample_advantages = [task_adv_weight * sample_advantage for sample_advantage in sample_advantages] prompt_repeat_k = len(group) for i in range(prompt_repeat_k): @@ -1207,15 +1196,9 @@ def _prepare_train_data( shifted_labels = [-100] * (len(prompt_ids) - 1) + response_labels shifted_labels_t = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0) - if opd_token_advantages is None: - advatnages_val = sample_advantages[i] - actual_advantages = [advatnages_val] * len(prompt_ids) + [ - 0.0 if mask == 0 else advatnages_val for mask in response_mask - ] - advantages_list.extend(actual_advantages[:-1]) - else: - actual_advantages = [0.0] * (len(prompt_ids) - 1) + opd_token_advantages[i].tolist() - advantages_list.extend(actual_advantages) + base_advantage = sample_advantages[i] + actual_advantages = [0.0 if label == -100 else base_advantage for label in shifted_labels] + advantages_list.extend(actual_advantages) assert len(input_ids) <= pack_max_length, f"{len(input_ids)} vs {pack_max_length}" training_tokens += len(input_ids) @@ -1240,6 +1223,12 @@ def _prepare_train_data( "advantage": actual_advantages, "rollout_logprobs": rollout_logprobs, } + if opd_config is not None: + response_teacher_logprobs = cast(list[float], group[i].teacher_logprobs) + data_dict["teacher_logprobs"] = torch.tensor( + [0.0] * (len(prompt_ids) - 1) + response_teacher_logprobs, + dtype=torch.float32, + ).unsqueeze(0) seq_ctx.rollout_routed_experts = group[i].routed_experts # n,layer*expert @@ -1396,6 +1385,8 @@ def _log_step( all_scalars.update({f"{k}": v for k, v in rank0_mismatch_metrics.items()}) all_scalars.update({"entropy/rollout": rank0_rollout_entropy}) all_scalars.update({"entropy/train": rank0_log_item["train_entropy"]}) + if "opd_reverse_kl" in rank0_log_item: + all_scalars["opd_reverse_kl"] = rank0_log_item["opd_reverse_kl"] for worker_idx, log_item in enumerate(train_info["workers_log_item"]): if not self._display_all_workers_log and worker_idx > 0: break From e90e9266c5c971a71c17fe1d73641b949de3e72d Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Tue, 28 Jul 2026 07:24:43 +0000 Subject: [PATCH 6/8] add eval in opd config --- .../rl_dapo_math_opd.py | 71 +++++++++++++++++-- ...run_pg_opd.sh => run_sampled_token_opd.sh} | 38 +--------- xtuner/v1/train/rl_trainer.py | 13 ++-- 3 files changed, 75 insertions(+), 47 deletions(-) rename recipe/on_policy_distillation/{run_pg_opd.sh => run_sampled_token_opd.sh} (73%) diff --git a/recipe/on_policy_distillation/rl_dapo_math_opd.py b/recipe/on_policy_distillation/rl_dapo_math_opd.py index 8e41b60439..70351650f3 100644 --- a/recipe/on_policy_distillation/rl_dapo_math_opd.py +++ b/recipe/on_policy_distillation/rl_dapo_math_opd.py @@ -1,6 +1,8 @@ import os from pathlib import Path +from transformers import AutoTokenizer + from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig from xtuner.v1.data_proto.rl_data import SampleParams from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig @@ -14,24 +16,27 @@ SyncProduceStrategyConfig, TaskSpecConfig, ) +from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.judger import DapoMathJudgerConfig from xtuner.v1.rl.loss import GRPOLossConfig from xtuner.v1.rl.on_policy_distillation import OPDConfig, OPDTeacherConfig from xtuner.v1.rl.replay_buffer import SyncReplayBufferConfig from xtuner.v1.rl.rollout.worker import RolloutConfig from xtuner.v1.rl.trainer import WorkerConfig -from xtuner.v1.rl.utils import AcceleratorResourcesConfig +from xtuner.v1.rl.utils import AcceleratorResourcesConfig, CPUResourcesConfig, get_eos_token from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig work_dir = os.environ["WORK_DIR"] model_path = os.environ["MODEL_PATH"] data_path = os.environ["DATA_PATH"] +eval_data_path = os.environ.get("EVAL_DATA_PATH", "") NNODE = int(os.environ.get("WORLD_SIZE", "1")) # basic settings experimental_name = "dapo_math_opd" total_train_steps = 300 -train_batch_size = 64 +train_batch_size = 16 prompt_repeat_k = 4 rollout_tp_size = 1 rollout_ep_size = 1 @@ -39,6 +44,9 @@ max_response_length = 16384 pack_max_length = max_prompt_length + max_response_length train_optimizer_steps = 1 +enable_evaluate = bool(eval_data_path) +evaluate_step = 20 +eval_prompt_repeat_k = 16 # 1. resources resources = AcceleratorResourcesConfig( @@ -56,7 +64,7 @@ dtype="bfloat16", tensor_parallel_size=rollout_tp_size, expert_parallel_size=rollout_ep_size, - gpu_memory_utilization=0.4, + gpu_memory_utilization=0.6, context_length=max_response_length + max_prompt_length, enable_return_routed_experts=False, rollout_max_batch_size_per_instance=2048, @@ -141,7 +149,55 @@ ), ) -# 5. pure on-policy distillation +# 5. evaluation +eval_agent_loop_manager_cfg = None +evaluator_config = None +if enable_evaluate: + eos_token_id = get_eos_token(model_path) + tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) + eos_token = tokenizer.convert_ids_to_tokens(eos_token_id) + eval_judger_config = DapoMathJudgerConfig( + judger_name="aime_math", + cpu_resources=CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1), + eos_token=eos_token, + enable_overlong_buffer=False, + ) + eval_dataset = DatasetConfig(name="aime", anno_path=eval_data_path, sample_ratio=1.0) + eval_dataset_cfg = [{"dataset": eval_dataset, "tokenize_fn": tokenizer_config}] + eval_dataloader_cfg = DataloaderConfig( + dataset_config_list=eval_dataset_cfg, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", + ) + eval_sampler_config = SamplerConfig( + dataloader_cfg=eval_dataloader_cfg, + prompt_repeat_k=eval_prompt_repeat_k, + ) + evaluation_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, + ) + eval_agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=evaluation_sample_params, + ) + eval_agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="eval_task", + agent_loop_config=eval_agent_loop_config, + judger_config=eval_judger_config, + sampler_config=eval_sampler_config, + ), + ) + evaluator_config = EvaluatorConfig() + +# 6. pure on-policy distillation opd_config = OPDConfig( mode="pg-opd", task_adv_weight=0.0, @@ -157,12 +213,15 @@ tokenizer_path=model_path, replay_buffer_config=SyncReplayBufferConfig(), agent_loop_manager_cfg=agent_loop_manager_cfg, + eval_agent_loop_manager_cfg=eval_agent_loop_manager_cfg, + evaluator_config=evaluator_config, load_from=model_path, train_batch_size=train_batch_size, advantage_estimator_config=GRPOAdvantageConfig(eps=1e-8), opd_config=opd_config, - enable_evaluate=False, - enable_initial_evaluate=False, + enable_evaluate=enable_evaluate, + enable_initial_evaluate=enable_evaluate, + evaluate_step=evaluate_step, total_train_steps=total_train_steps, work_dir=work_dir, seed=1234, diff --git a/recipe/on_policy_distillation/run_pg_opd.sh b/recipe/on_policy_distillation/run_sampled_token_opd.sh similarity index 73% rename from recipe/on_policy_distillation/run_pg_opd.sh rename to recipe/on_policy_distillation/run_sampled_token_opd.sh index b4f8aae7ee..62b2e0c009 100755 --- a/recipe/on_policy_distillation/run_pg_opd.sh +++ b/recipe/on_policy_distillation/run_sampled_token_opd.sh @@ -5,15 +5,9 @@ set -euo pipefail SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../.." && pwd) -if (( $# != 0 )); then - echo "This script does not accept positional arguments." >&2 - echo "Set STUDENT_MODEL_PATH, TEACHER_MODEL_PATH, and DATA_PATH before running it." >&2 - exit 2 -fi - -: "${STUDENT_MODEL_PATH:?STUDENT_MODEL_PATH is required}" -: "${TEACHER_MODEL_PATH:?TEACHER_MODEL_PATH is required}" -: "${DATA_PATH:?DATA_PATH is required}" +STUDENT_MODEL_PATH="$1" +TEACHER_MODEL_PATH="$2" +DATA_PATH="$3" STUDENT_CUDA_VISIBLE_DEVICES="0,1,2,3" TEACHER_CUDA_VISIBLE_DEVICES="7" @@ -28,34 +22,8 @@ WORK_DIR="${REPO_ROOT}/work_dirs/dapo_math_opd" TEACHER_ENDPOINT="http://${TEACHER_HOST}:${TEACHER_PORT}" TEACHER_LOG_FILE="${WORK_DIR}/teacher.log" -if [[ ! -d "${STUDENT_MODEL_PATH}" ]]; then - echo "Student model directory does not exist: ${STUDENT_MODEL_PATH}" >&2 - exit 1 -fi -if [[ ! -d "${TEACHER_MODEL_PATH}" ]]; then - echo "Teacher model directory does not exist: ${TEACHER_MODEL_PATH}" >&2 - exit 1 -fi -if [[ ! -f "${DATA_PATH}" ]]; then - echo "Training data file does not exist: ${DATA_PATH}" >&2 - exit 1 -fi - -for required_command in python curl ray setsid; do - if ! command -v "${required_command}" >/dev/null 2>&1; then - echo "Required command is not available: ${required_command}" >&2 - exit 1 - fi -done - -python -c "import sglang" >/dev/null mkdir -p "${WORK_DIR}" -if curl -sf --max-time 2 "${TEACHER_ENDPOINT}/health_generate" >/dev/null; then - echo "A teacher service is already running at ${TEACHER_ENDPOINT}" >&2 - exit 1 -fi - TEACHER_PID="" TRAINING_STARTED=0 diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 6e7f003e3b..89e4f89404 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -1424,8 +1424,9 @@ def _save_trajectories(self, data_groups: list[list[RolloutState]], save_path: P if not is_valid_for_training(group, self.logger): continue for data in group: - assert data.reward is not None - rewards.append(data.reward["score"]) + reward = data.reward.get("score") if data.reward is not None else None + if reward is not None: + rewards.append(reward) response_ids = self._get_trajectory_response_ids(data) response = data.response if response is None and response_ids: @@ -1446,7 +1447,7 @@ def _save_trajectories(self, data_groups: list[list[RolloutState]], save_path: P "prompt": data.message, "label": ground_truth, "response": response, - "reward": data.reward["score"], + "reward": reward, "prompt_len": data.num_tokens, "response_len": response_len, "reward_payload": data.reward, @@ -1471,14 +1472,14 @@ def _save_trajectories(self, data_groups: list[list[RolloutState]], save_path: P with open(save_path, "w", encoding="utf-8") as f: summary = { "reward_mean": rewards_tensor.mean().item(), - "reward_std": rewards_tensor.std().item(), + "reward_std": rewards_tensor.std(unbiased=False).item(), "reward_max": rewards_tensor.max().item(), "reward_min": rewards_tensor.min().item(), "response_len_mean": response_lens.mean().item(), - "response_len_std": response_lens.std().item(), + "response_len_std": response_lens.std(unbiased=False).item(), "response_len_max": response_lens.max().item(), "response_len_min": response_lens.min().item(), - "total_len": len(rewards), + "total_len": len(trajectory_items), } json.dump(summary, f, ensure_ascii=False, separators=(",", ":")) f.write("\n") From 1ed9955ec3533f54e9843b0efb72b43dc6d08de9 Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Tue, 28 Jul 2026 09:46:26 +0000 Subject: [PATCH 7/8] add lmdeploy backend --- .../run_sampled_token_opd.sh | 41 ++++-- .../start_pg_opd_teacher.sh | 40 +++++- tests/rl/test_on_policy_distillation.py | 55 ++++---- xtuner/v1/rl/on_policy_distillation.py | 121 +++++++++++++++--- 4 files changed, 198 insertions(+), 59 deletions(-) diff --git a/recipe/on_policy_distillation/run_sampled_token_opd.sh b/recipe/on_policy_distillation/run_sampled_token_opd.sh index 62b2e0c009..a2a5e28e0b 100755 --- a/recipe/on_policy_distillation/run_sampled_token_opd.sh +++ b/recipe/on_policy_distillation/run_sampled_token_opd.sh @@ -18,6 +18,27 @@ TEACHER_CHUNKED_PREFILL_SIZE="4096" TEACHER_GPU_MEMORY_UTILIZATION="0.6" TEACHER_STARTUP_TIMEOUT_S="1200" +USE_SGLANG=${XTUNER_USE_SGLANG:-0} +USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} +USE_VLLM=${XTUNER_USE_VLLM:-0} + +if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="sglang" + TEACHER_HEALTH_PATH="health_generate" + TEACHER_MODEL_INFO_PATH="get_model_info" +elif [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="lmdeploy" + TEACHER_HEALTH_PATH="health" + TEACHER_MODEL_INFO_PATH="v1/models" +else + echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 + exit 1 +fi + +export XTUNER_USE_SGLANG="${USE_SGLANG}" +export XTUNER_USE_LMDEPLOY="${USE_LMDEPLOY}" +export XTUNER_USE_VLLM="${USE_VLLM}" + WORK_DIR="${REPO_ROOT}/work_dirs/dapo_math_opd" TEACHER_ENDPOINT="http://${TEACHER_HOST}:${TEACHER_PORT}" TEACHER_LOG_FILE="${WORK_DIR}/teacher.log" @@ -63,7 +84,7 @@ wait_for_teacher() { tail -n 50 "${TEACHER_LOG_FILE}" >&2 || true return 1 fi - if curl -sf --max-time 2 "${TEACHER_ENDPOINT}/health_generate" >/dev/null; then + if curl -sf --max-time 2 "${TEACHER_ENDPOINT}/${TEACHER_HEALTH_PATH}" >/dev/null; then return 0 fi echo "Waiting for teacher service at ${TEACHER_ENDPOINT}..." @@ -80,24 +101,24 @@ trap "exit 130" INT trap "exit 143" TERM echo "Starting teacher model: ${TEACHER_MODEL_PATH}" +echo "Teacher backend: ${OPD_BACKEND}" echo "Teacher GPUs: ${TEACHER_CUDA_VISIBLE_DEVICES}" echo "Teacher log: ${TEACHER_LOG_FILE}" setsid env \ CUDA_VISIBLE_DEVICES="${TEACHER_CUDA_VISIBLE_DEVICES}" \ PYTHONUNBUFFERED=1 \ - python -m sglang.launch_server \ - --model-path "${TEACHER_MODEL_PATH}" \ - --host "${TEACHER_HOST}" \ - --port "${TEACHER_PORT}" \ - --tp "${TEACHER_TP_SIZE}" \ - --chunked-prefill-size "${TEACHER_CHUNKED_PREFILL_SIZE}" \ - --mem-fraction-static "${TEACHER_GPU_MEMORY_UTILIZATION}" \ + TEACHER_HOST="${TEACHER_HOST}" \ + TEACHER_PORT="${TEACHER_PORT}" \ + TEACHER_TP_SIZE="${TEACHER_TP_SIZE}" \ + TEACHER_CHUNKED_PREFILL_SIZE="${TEACHER_CHUNKED_PREFILL_SIZE}" \ + TEACHER_GPU_MEMORY_UTILIZATION="${TEACHER_GPU_MEMORY_UTILIZATION}" \ + bash "${SCRIPT_DIR}/start_pg_opd_teacher.sh" "${TEACHER_MODEL_PATH}" \ >"${TEACHER_LOG_FILE}" 2>&1 & TEACHER_PID=$! wait_for_teacher -curl -sS --max-time 10 "${TEACHER_ENDPOINT}/get_model_info" +curl -sS --max-time 10 "${TEACHER_ENDPOINT}/${TEACHER_MODEL_INFO_PATH}" echo echo "Teacher service is ready at ${TEACHER_ENDPOINT}" echo "Starting Pure PG-OPD training with student GPUs: ${STUDENT_CUDA_VISIBLE_DEVICES}" @@ -110,6 +131,6 @@ TRAINING_STARTED=1 CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES}" \ bash -o pipefail examples/v1/scripts/run_rl.sh \ recipe/on_policy_distillation/rl_dapo_math_opd.py \ - sglang \ + "${OPD_BACKEND}" \ "${STUDENT_MODEL_PATH}" \ "${DATA_PATH}" diff --git a/recipe/on_policy_distillation/start_pg_opd_teacher.sh b/recipe/on_policy_distillation/start_pg_opd_teacher.sh index 575e0210f2..477bffcc14 100755 --- a/recipe/on_policy_distillation/start_pg_opd_teacher.sh +++ b/recipe/on_policy_distillation/start_pg_opd_teacher.sh @@ -2,6 +2,9 @@ set -euo pipefail +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../.." && pwd) + TEACHER_MODEL_PATH=${1:?"Usage: $0 TEACHER_MODEL_PATH"} TEACHER_HOST=${TEACHER_HOST:-127.0.0.1} TEACHER_PORT=${TEACHER_PORT:-13141} @@ -9,11 +12,34 @@ TEACHER_TP_SIZE=${TEACHER_TP_SIZE:-1} TEACHER_CHUNKED_PREFILL_SIZE=${TEACHER_CHUNKED_PREFILL_SIZE:-4096} TEACHER_GPU_MEMORY_UTILIZATION=${TEACHER_GPU_MEMORY_UTILIZATION:-0.6} PYTHON_EXECUTABLE=${PYTHON_EXECUTABLE:-python} +LMDEPLOY_PATH=${LMDEPLOY_PATH:-"${REPO_ROOT}/work_dirs/lmdeploy"} + +USE_SGLANG=${XTUNER_USE_SGLANG:-0} +USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} +USE_VLLM=${XTUNER_USE_VLLM:-0} + +if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then + exec "${PYTHON_EXECUTABLE}" -m sglang.launch_server \ + --model-path "${TEACHER_MODEL_PATH}" \ + --host "${TEACHER_HOST}" \ + --port "${TEACHER_PORT}" \ + --tp "${TEACHER_TP_SIZE}" \ + --chunked-prefill-size "${TEACHER_CHUNKED_PREFILL_SIZE}" \ + --mem-fraction-static "${TEACHER_GPU_MEMORY_UTILIZATION}" +fi + +if [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then + export PYTHONPATH="${LMDEPLOY_PATH}${PYTHONPATH:+:${PYTHONPATH}}" + exec "${PYTHON_EXECUTABLE}" -m lmdeploy serve api_server "${TEACHER_MODEL_PATH}" \ + --backend pytorch \ + --role Hybrid \ + --logprobs-mode raw_logprobs \ + --server-name "${TEACHER_HOST}" \ + --server-port "${TEACHER_PORT}" \ + --tp "${TEACHER_TP_SIZE}" \ + --max-prefill-token-num "${TEACHER_CHUNKED_PREFILL_SIZE}" \ + --cache-max-entry-count "${TEACHER_GPU_MEMORY_UTILIZATION}" +fi -exec "${PYTHON_EXECUTABLE}" -m sglang.launch_server \ - --model-path "${TEACHER_MODEL_PATH}" \ - --host "${TEACHER_HOST}" \ - --port "${TEACHER_PORT}" \ - --tp "${TEACHER_TP_SIZE}" \ - --chunked-prefill-size "${TEACHER_CHUNKED_PREFILL_SIZE}" \ - --mem-fraction-static "${TEACHER_GPU_MEMORY_UTILIZATION}" +echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 +exit 1 diff --git a/tests/rl/test_on_policy_distillation.py b/tests/rl/test_on_policy_distillation.py index df87baa047..5f23c22069 100644 --- a/tests/rl/test_on_policy_distillation.py +++ b/tests/rl/test_on_policy_distillation.py @@ -29,15 +29,16 @@ TRAINER_CONFIG_PATH = REPO_ROOT / "recipe/on_policy_distillation/rl_dapo_math_opd.py" -def _wait_for_teacher(process: subprocess.Popen, endpoint: str) -> None: +def _wait_for_teacher(process: subprocess.Popen, endpoint: str, backend: str) -> None: deadline = time.monotonic() + TEACHER_STARTUP_TIMEOUT_S + health_path = "health_generate" if backend == "sglang" else "health" with httpx.Client(timeout=1.0, trust_env=False) as client: while time.monotonic() < deadline: return_code = process.poll() if return_code is not None: raise RuntimeError(f"Teacher process exited during startup with code {return_code}") try: - if client.get(f"{endpoint}/health_generate").status_code == 200: + if client.get(f"{endpoint}/{health_path}").status_code == 200: return except httpx.RequestError: pass @@ -246,7 +247,7 @@ def fit(self, data_batches, rollout_idx): def capture_opd_result(loss_ctx, *, config): nonlocal captured_batch_index - apply_opd(loss_ctx, config=config) + reverse_kl_sum = apply_opd(loss_ctx, config=config) captured_batches[captured_batch_index]["old_log_probs"] = ( loss_ctx.loss_kwargs.old_logprobs.detach().cpu().reshape(-1) ) @@ -254,6 +255,7 @@ def capture_opd_result(loss_ctx, *, config): loss_ctx.loss_kwargs.advantages.detach().cpu().reshape(-1) ) captured_batch_index += 1 + return reverse_kl_sum with mock_patch.object(worker_module, "apply_opd_kl_to_advantages", capture_opd_result): worker_log_item = TrainingWorker.fit(self, data_batches, rollout_idx) @@ -483,6 +485,7 @@ def setUpClass(cls) -> None: cls.samples = baseline["samples"] port = find_free_ports()[0] cls.teacher_endpoint = f"http://127.0.0.1:{port}" + cls.teacher_backend = TeacherLogprobClient._resolve_backend_from_env() teacher_env = os.environ.copy() teacher_env["TEACHER_HOST"] = "127.0.0.1" teacher_env["TEACHER_PORT"] = str(port) @@ -493,7 +496,7 @@ def setUpClass(cls) -> None: start_new_session=True, ) cls.addClassCleanup(cls._stop_teacher) - _wait_for_teacher(cls.teacher_process, cls.teacher_endpoint) + _wait_for_teacher(cls.teacher_process, cls.teacher_endpoint, cls.teacher_backend) @classmethod def _stop_teacher(cls) -> None: @@ -506,7 +509,17 @@ def _stop_teacher(cls) -> None: os.killpg(cls.teacher_process.pid, signal.SIGKILL) cls.teacher_process.wait() - async def test_compute_logprobs_matches_baseline(self) -> None: + @unittest.skipUnless(os.getenv("XTUNER_USE_SGLANG", "0") == "1", "XTUNER_USE_SGLANG=1 is required") + async def test_compute_logprobs_with_sglang_matches_baseline(self) -> None: + self.assertEqual(self.teacher_backend, "sglang") + await self._assert_compute_logprobs_matches_baseline() + + @unittest.skipUnless(os.getenv("XTUNER_USE_LMDEPLOY", "0") == "1", "XTUNER_USE_LMDEPLOY=1 is required") + async def test_compute_logprobs_with_lmdeploy_matches_baseline(self) -> None: + self.assertEqual(self.teacher_backend, "lmdeploy") + await self._assert_compute_logprobs_matches_baseline() + + async def _assert_compute_logprobs_matches_baseline(self) -> None: client = TeacherLogprobClient(OPDTeacherConfig(name="teacher", endpoint=self.teacher_endpoint)) self.addAsyncCleanup(client._client.aclose) @@ -531,27 +544,17 @@ async def test_compute_logprobs_matches_baseline(self) -> None: self.assertEqual(result.status, Status.COMPLETED, result.error_msg) self.assertEqual(result.teacher_tokens, response_ids) actual_logprobs = torch.tensor(result.teacher_logprobs, dtype=torch.float32) - try: - torch.testing.assert_close( - actual_logprobs, - expected_logprobs, - rtol=1e-5, - atol=1e-5, - ) - except AssertionError as error: - mismatch_indices = torch.nonzero( - ~torch.isclose(actual_logprobs, expected_logprobs, rtol=1e-5, atol=1e-5) - ).flatten() - mismatch_values = "\n".join( - ( - f"index={index}: " - f"actual={actual_logprobs[index].item()!r}, " - f"expected={expected_logprobs[index].item()!r}, " - f"abs_diff={abs(actual_logprobs[index] - expected_logprobs[index]).item()!r}" - ) - for index in mismatch_indices.tolist() - ) - raise AssertionError(f"{error}\n\nMismatched logprobs:\n{mismatch_values}") from None + self.assertEqual(actual_logprobs.shape, expected_logprobs.shape) + absolute_errors = torch.abs(actual_logprobs - expected_logprobs) + mae = absolute_errors.mean().item() + self.assertLessEqual( + mae, + 0.05, + ( + f"Teacher logprob MAE {mae:.8f} exceeds threshold 0.05; " + f"max_abs_error={absolute_errors.max().item():.8f}" + ), + ) if __name__ == "__main__": diff --git a/xtuner/v1/rl/on_policy_distillation.py b/xtuner/v1/rl/on_policy_distillation.py index 62d8961f95..9ca18a2f30 100644 --- a/xtuner/v1/rl/on_policy_distillation.py +++ b/xtuner/v1/rl/on_policy_distillation.py @@ -2,6 +2,7 @@ import asyncio import math +import os import time from typing import Any, Literal, cast @@ -66,11 +67,12 @@ def validate_opd_sample_params(sample_params: SampleParams) -> None: class TeacherLogprobClient: - """Asynchronous SGLang teacher client scoped to one AgentLoop.""" + """Asynchronous teacher client scoped to one AgentLoop.""" def __init__(self, config: OPDTeacherConfig) -> None: self.config = config self.name = config.name + self.backend = self._resolve_backend_from_env() self.url = f"{config.endpoint.rstrip('/')}/generate" self._semaphore = asyncio.Semaphore(config.max_concurrency) @@ -84,18 +86,11 @@ async def compute_logprobs(self, state: RolloutState) -> RolloutState: try: prompt_ids = cast(list[int], state.prompt_ids) response_ids = cast(list[int], state.response_ids) - payload = { - "input_ids": prompt_ids + response_ids, - "sampling_params": { - "max_new_tokens": 0, - "temperature": 0, - "skip_special_tokens": False, - }, - "return_logprob": True, - "logprob_start_len": 0, - "top_logprobs_num": 0, - "stream": False, - } + if not prompt_ids or not response_ids: + state.status = Status.FAILED + state.error_msg = f"Teacher {self.name!r} scoring requires non-empty prompt_ids and response_ids" + return state + payload = self._construct_payload(prompt_ids, response_ids) retries = 0 while True: @@ -103,7 +98,7 @@ async def compute_logprobs(self, state: RolloutState) -> RolloutState: async with self._semaphore: response = await self._client.post(self.url, json=payload) response.raise_for_status() - teacher_tokens, teacher_logprobs = self._parse_response(response, response_ids) + teacher_tokens, teacher_logprobs = self._parse_response(response, prompt_ids, response_ids) state.teacher_tokens = teacher_tokens state.teacher_logprobs = teacher_logprobs return state @@ -117,17 +112,111 @@ async def compute_logprobs(self, state: RolloutState) -> RolloutState: finally: state.extra_fields["teacher_score_time_s"] = time.perf_counter() - start + @staticmethod + def _resolve_backend_from_env() -> Literal["sglang", "lmdeploy"]: + use_sglang = os.environ.get("XTUNER_USE_SGLANG", "0") == "1" + use_lmdeploy = os.environ.get("XTUNER_USE_LMDEPLOY", "0") == "1" + use_vllm = os.environ.get("XTUNER_USE_VLLM", "0") == "1" + + if use_vllm: + raise RuntimeError("TeacherLogprobClient supports only SGLang or LMDeploy, not vLLM") + if use_sglang == use_lmdeploy: + raise RuntimeError("Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1") + return "sglang" if use_sglang else "lmdeploy" + + def _construct_payload(self, prompt_ids: list[int], response_ids: list[int]) -> dict[str, Any]: + if self.backend == "sglang": + return self._construct_sglang_payload(prompt_ids, response_ids) + if self.backend == "lmdeploy": + return self._construct_lmdeploy_payload(prompt_ids, response_ids) + raise RuntimeError(f"Unsupported teacher backend: {self.backend}") + + @staticmethod + def _construct_sglang_payload(prompt_ids: list[int], response_ids: list[int]) -> dict[str, Any]: + return { + "input_ids": prompt_ids + response_ids, + "sampling_params": { + "max_new_tokens": 0, + "temperature": 0, + "skip_special_tokens": False, + }, + "return_logprob": True, + "logprob_start_len": 0, + "top_logprobs_num": 0, + "stream": False, + } + + @staticmethod + def _construct_lmdeploy_payload(prompt_ids: list[int], response_ids: list[int]) -> dict[str, Any]: + return { + "input_ids": prompt_ids + response_ids, + "return_logprob": True, + "logprob_start_len": 0, + "max_tokens": 0, + "stream": False, + } + def _parse_response( self, response: httpx.Response, + prompt_ids: list[int], response_ids: list[int], ) -> tuple[list[int], list[float]]: + if self.backend == "sglang": + response_logprobs = self._parse_sglang_response(response, prompt_ids, response_ids) + else: + response_logprobs = self._parse_lmdeploy_response(response, prompt_ids, response_ids) + return self._validate_response_logprobs(response_logprobs, response_ids) + + @staticmethod + def _parse_sglang_response( + response: httpx.Response, + prompt_ids: list[int], + response_ids: list[int], + ) -> list[Any]: + raw_logprobs = TeacherLogprobClient._get_input_token_logprobs(response) + expected_length = len(prompt_ids) + len(response_ids) + if len(raw_logprobs) != expected_length: + raise ValueError( + "SGLang teacher logprob length mismatch: " + f"expected {expected_length} rows for the full input, got {len(raw_logprobs)}" + ) + return raw_logprobs[-len(response_ids) :] + + @staticmethod + def _parse_lmdeploy_response( + response: httpx.Response, + prompt_ids: list[int], + response_ids: list[int], + ) -> list[Any]: + raw_logprobs = TeacherLogprobClient._get_input_token_logprobs(response) + expected_length = len(prompt_ids) + len(response_ids) - 1 + if len(raw_logprobs) != expected_length: + raise ValueError( + "LMDeploy teacher logprob length mismatch: " + f"expected {expected_length} rows after the boundary token, got {len(raw_logprobs)}" + ) + return raw_logprobs[-len(response_ids) :] + + @staticmethod + def _get_input_token_logprobs(response: httpx.Response) -> list[Any]: try: raw_logprobs = response.json()["meta_info"]["input_token_logprobs"] - response_logprobs = raw_logprobs[-len(response_ids) :] + except (KeyError, TypeError, ValueError) as exc: + raise ValueError("Invalid teacher response") from exc + if not isinstance(raw_logprobs, list): + raise ValueError("Invalid teacher response") + return raw_logprobs + + @staticmethod + def _validate_response_logprobs( + response_logprobs: list[Any], + response_ids: list[int], + ) -> tuple[list[int], list[float]]: + try: teacher_tokens = [item[1] for item in response_logprobs] teacher_logprobs = [float(item[0]) for item in response_logprobs] - except (KeyError, TypeError, IndexError) as exc: + except (TypeError, IndexError, ValueError) as exc: raise ValueError("Invalid teacher response") from exc if len(teacher_logprobs) != len(response_ids): From 8d0e07ce53bad6d1aa2fee78738b91175b68b69a Mon Sep 17 00:00:00 2001 From: YanhuiDua Date: Wed, 29 Jul 2026 12:50:23 +0000 Subject: [PATCH 8/8] add OPDTeacherLaunchConfig, multi-teacher opd config and scripts --- .../build_teacher_server_commands.py | 196 ++++++++++ .../config/rl_dapo_math_mopd.py | 335 ++++++++++++++++++ .../{ => config}/rl_dapo_math_opd.py | 25 +- .../run_sampled_token_opd.sh | 136 ------- .../scripts/launch_teacher_utils.sh | 247 +++++++++++++ .../scripts/run_sampled_token_mopd.sh | 75 ++++ .../scripts/run_sampled_token_opd.sh | 72 ++++ .../start_pg_opd_teacher.sh | 45 --- tests/rl/test_on_policy_distillation.py | 29 +- xtuner/v1/rl/on_policy_distillation.py | 23 ++ 10 files changed, 993 insertions(+), 190 deletions(-) create mode 100644 recipe/on_policy_distillation/build_teacher_server_commands.py create mode 100644 recipe/on_policy_distillation/config/rl_dapo_math_mopd.py rename recipe/on_policy_distillation/{ => config}/rl_dapo_math_opd.py (92%) delete mode 100755 recipe/on_policy_distillation/run_sampled_token_opd.sh create mode 100644 recipe/on_policy_distillation/scripts/launch_teacher_utils.sh create mode 100755 recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh create mode 100755 recipe/on_policy_distillation/scripts/run_sampled_token_opd.sh delete mode 100755 recipe/on_policy_distillation/start_pg_opd_teacher.sh diff --git a/recipe/on_policy_distillation/build_teacher_server_commands.py b/recipe/on_policy_distillation/build_teacher_server_commands.py new file mode 100644 index 0000000000..87ddc68e4e --- /dev/null +++ b/recipe/on_policy_distillation/build_teacher_server_commands.py @@ -0,0 +1,196 @@ +"""Build executable Teacher server commands from an OPD config. + +The NUL-delimited output starts with the Teacher count. Each Teacher record +contains its name, safe name, endpoint, health URL, model-info URL, +command-argument count, and executable command arguments. An externally +managed Teacher has a command-argument count of zero. +""" + +import argparse +import re +import sys +from pathlib import Path +from typing import Literal +from urllib.parse import urlparse + + +REPO_ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(REPO_ROOT)) + +from xtuner.v1.rl.on_policy_distillation import OPDTeacherConfig # noqa: E402 +from xtuner.v1.utils.config import Config # noqa: E402 + + +def build_teacher_server_command( + teacher: OPDTeacherConfig, + backend: Literal["sglang", "lmdeploy"], +) -> list[str]: + """Build one executable Teacher server command.""" + if backend == "sglang": + return _build_sglang_command(teacher) + return _build_lmdeploy_command(teacher) + + +def build_teacher_server_commands( + config_path: str, + backend: Literal["sglang", "lmdeploy"], +) -> list[list[str]]: + """Build executable Teacher server command records. + + Args: + config_path (str): Path to the XTuner Python config containing + ``opd_config``. + backend (Literal["sglang", "lmdeploy"]): Teacher serving backend. + + Returns: + list[list[str]]: Teacher metadata and executable argv records. + """ + config = Config.fromfile(config_path) + health_path, model_info_path = _get_teacher_server_paths(backend) + records: list[list[str]] = [] + for teacher in config.opd_config.teachers: + command = build_teacher_server_command(teacher, backend) + safe_name = re.sub(r"[^A-Za-z0-9._-]+", "_", teacher.name).strip("_") or "teacher" + endpoint = teacher.endpoint.rstrip("/") + records.append( + [ + teacher.name, + safe_name, + endpoint, + f"{endpoint}/{health_path}", + f"{endpoint}/{model_info_path}", + str(len(command)), + *command, + ] + ) + return records + + +def _get_teacher_server_paths( + backend: Literal["sglang", "lmdeploy"], +) -> tuple[str, str]: + if backend == "sglang": + return "health_generate", "get_model_info" + return "health", "v1/models" + + +def _parse_teacher_endpoint(teacher: OPDTeacherConfig) -> tuple[str, int]: + parsed_endpoint = urlparse(teacher.endpoint) + if parsed_endpoint.scheme not in {"http", "https"}: + raise ValueError(f"Teacher {teacher.name!r} endpoint must use http or https: {teacher.endpoint}") + if parsed_endpoint.hostname is None or parsed_endpoint.port is None: + raise ValueError( + f"Teacher {teacher.name!r} endpoint must contain an explicit host and port: {teacher.endpoint}" + ) + if parsed_endpoint.path not in {"", "/"} or parsed_endpoint.params or parsed_endpoint.query: + raise ValueError( + f"Teacher {teacher.name!r} endpoint must be a server base URL without a path or query: {teacher.endpoint}" + ) + return parsed_endpoint.hostname, parsed_endpoint.port + + +def _build_sglang_command(teacher: OPDTeacherConfig) -> list[str]: + host, port = _parse_teacher_endpoint(teacher) + config = teacher.launch_config + if config is None: + return [] + + tensor_parallel_size = config.tensor_parallel_size + if config.expert_parallel_size > 1: + tensor_parallel_size = config.expert_parallel_size + + command = [ + "env", + f"CUDA_VISIBLE_DEVICES={config.cuda_visible_devices}", + sys.executable, + "-m", + "sglang.launch_server", + "--model-path", + str(config.model_path), + "--host", + host, + "--port", + str(port), + "--dtype", + config.dtype, + "--tp", + str(tensor_parallel_size), + "--ep", + str(config.expert_parallel_size), + "--mem-fraction-static", + str(config.gpu_memory_utilization), + ] + if config.context_length is not None: + command.extend(["--context-length", str(config.context_length)]) + if config.max_batch_size is not None: + command.extend(["--max-running-requests", str(config.max_batch_size)]) + if config.chunked_prefill_size is not None: + command.extend(["--chunked-prefill-size", str(config.chunked_prefill_size)]) + return command + + +def _build_lmdeploy_command(teacher: OPDTeacherConfig) -> list[str]: + host, port = _parse_teacher_endpoint(teacher) + config = teacher.launch_config + if config is None: + return [] + + data_parallel_size = config.expert_parallel_size if config.expert_parallel_size > 1 else 1 + command = [ + "env", + f"CUDA_VISIBLE_DEVICES={config.cuda_visible_devices}", + sys.executable, + "-m", + "lmdeploy", + "serve", + "api_server", + str(config.model_path), + "--backend", + "pytorch", + "--role", + "Hybrid", + "--logprobs-mode", + "raw_logprobs", + "--server-name", + host, + "--server-port", + str(port), + "--dtype", + config.dtype, + "--tp", + str(config.tensor_parallel_size), + "--ep", + str(config.expert_parallel_size), + "--dp", + str(data_parallel_size), + "--cache-max-entry-count", + str(config.gpu_memory_utilization), + ] + if config.context_length is not None: + command.extend(["--session-len", str(config.context_length)]) + if config.max_batch_size is not None: + command.extend(["--max-batch-size", str(config.max_batch_size)]) + if config.max_prefill_token_num is not None: + command.extend(["--max-prefill-token-num", str(config.max_prefill_token_num)]) + return command + + +def _write_teacher_records(records: list[list[str]]) -> None: + fields = [str(len(records))] + for record in records: + fields.extend(record) + + payload = "\0".join(fields) + "\0" + sys.stdout.buffer.write(payload.encode("utf-8")) + + +def _main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("config_path") + parser.add_argument("backend", choices=("sglang", "lmdeploy")) + args = parser.parse_args() + _write_teacher_records(build_teacher_server_commands(args.config_path, args.backend)) + + +if __name__ == "__main__": + _main() diff --git a/recipe/on_policy_distillation/config/rl_dapo_math_mopd.py b/recipe/on_policy_distillation/config/rl_dapo_math_mopd.py new file mode 100644 index 0000000000..a00ed68d80 --- /dev/null +++ b/recipe/on_policy_distillation/config/rl_dapo_math_mopd.py @@ -0,0 +1,335 @@ +import json +import os +from pathlib import Path + +from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig +from xtuner.v1.data_proto.rl_data import SampleParams +from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig +from xtuner.v1.datasets.rl_tokenize_fn import RLQwen3VLTokenizeFnConfig +from xtuner.v1.model import get_model_config_from_hf +from xtuner.v1.rl.advantage import GRPOAdvantageConfig +from xtuner.v1.rl.agent_loop import SingleTurnAgentLoopConfig +from xtuner.v1.rl.agent_loop_manager import ( + AgentLoopManagerConfig, + SamplerConfig, + SyncProduceStrategyConfig, + TaskSpecConfig, +) +from xtuner.v1.rl.evaluator import EvaluatorConfig +from xtuner.v1.rl.judger import ComposedJudgerConfig, GEO3KJudgerConfig, GSM8KJudgerConfig +from xtuner.v1.rl.loss import GRPOLossConfig +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + OPDTeacherConfig, + OPDTeacherLaunchConfig, +) +from xtuner.v1.rl.replay_buffer import SyncReplayBufferConfig +from xtuner.v1.rl.rollout.worker import RolloutConfig +from xtuner.v1.rl.trainer import WorkerConfig +from xtuner.v1.rl.utils import ( + AcceleratorResourcesConfig, + CPUResourcesConfig, +) +from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig + + +work_dir = os.environ["WORK_DIR"] +model_path = os.environ["MODEL_PATH"] +gsm8k_teacher_model_path = os.environ["GSM8K_TEACHER_MODEL_PATH"] +geo3k_teacher_model_path = os.environ["GEO3K_TEACHER_MODEL_PATH"] +meta_data_path = os.environ["DATA_PATH"] +eval_meta_data_path = os.environ.get("EVAL_DATA_PATH", "") +NNODE = int(os.environ.get("WORLD_SIZE", "1")) + + +def _as_list(value): + return value if isinstance(value, list) else [value] + + +# Training shape aligned with verl PR #6051: +# examples/on_policy_distillation_trainer/run_qwen3_mopd_gsm8k_geo3k.sh. +# Teacher roles and model families follow the GSM8K/Geo3K experiment. +experimental_name = "dapo_math_mopd" +total_epochs = 15 +train_batch_size = 128 +prompt_repeat_k = 1 +rollout_tp_size = 1 +rollout_ep_size = 1 +max_prompt_length = 1024 +max_response_length = 2048 +pack_max_length = max_prompt_length + max_response_length +max_num_tokens = pack_max_length +train_optimizer_steps = 1 +enable_evaluate = bool(eval_meta_data_path) +evaluate_step = 5 +eval_prompt_repeat_k = 1 +checkpoint_interval = 200 + +# 1. resources: two colocated Student workers, plus one GPU per Teacher. +resources = AcceleratorResourcesConfig( + accelerator="GPU", + num_workers=2 * NNODE, + num_cpus_per_worker=12, + cpu_memory_per_worker=16 * 1024**3, # 16 GB +) + +# 2. rollout +rollout_config = RolloutConfig( + env=experimental_name, + device=resources.accelerator, + model_path=model_path, + dtype="bfloat16", + tensor_parallel_size=rollout_tp_size, + expert_parallel_size=rollout_ep_size, + gpu_memory_utilization=0.6, + context_length=max_response_length + max_prompt_length, + enable_return_routed_experts=False, + rollout_max_batch_size_per_instance=2048, +) + +# 3. train worker +lr_cfg = LRConfig(lr_type="constant", warmup_ratio=0, lr_min=1e-6) +fsdp_cfg = FSDPConfig( + torch_compile=False, + cpu_offload=False, + ep_size=1, + reduce_dtype="float32", +) +model_cfg = get_model_config_from_hf(Path(model_path)) +if hasattr(model_cfg, "balancing_loss_cfg"): + model_cfg.balancing_loss_cfg = None +if hasattr(model_cfg, "z_loss_cfg"): + model_cfg.z_loss_cfg = None +optim_cfg = AdamWConfig(lr=1e-6, foreach=False, weight_decay=0.1, betas=(0.9, 0.98)) +loss_cfg = GRPOLossConfig( + policy_loss_cfg={ + "cliprange_high": 0.2, + "cliprange_low": 0.2, + "loss_type": "vanilla", + "clip_ratio_c": 10.0, + "log_prob_diff_min": -20.0, + "log_prob_diff_max": 20.0, + }, + ignore_idx=-100, + use_kl_loss=False, + kl_loss_coef=0.0, + kl_loss_type="low_var_kl", + mode="chunk", + chunk_size=512, +) +train_worker_cfg = WorkerConfig( + model_cfg=model_cfg, + load_from=model_path, + optim_cfg=optim_cfg, + loss_cfg=loss_cfg, + lr_cfg=lr_cfg, + fsdp_cfg=fsdp_cfg, + sp_size=1, + optimizer_steps=train_optimizer_steps, + pack_max_length=pack_max_length, +) + +# 4. train agent loop manager +with open(meta_data_path, "r", encoding="utf-8") as f: + ds_collections = json.load(f) + +train_dataset_cfg = [] +for name, data in ds_collections.items(): + annotations = _as_list(data["annotation"]) + for annotation in annotations: + train_dataset_cfg.append( + { + "dataset": DatasetConfig( + name=name, + anno_path=annotation, + media_root=data.get("media_root", ""), + sample_ratio=data.get("sample_ratio", 1.0), + class_name="VLMJsonlDataset", + ), + "tokenize_fn": RLQwen3VLTokenizeFnConfig( + processor_path=model_path, + max_length=max_prompt_length, + system_message=data.get("system_message", None), + chat_template="qwen3-vl", + add_generation_prompt=True, + enable_thinking=True, + ), + } + ) + +dataloader_cfg = DataloaderConfig( + dataset_config_list=train_dataset_cfg, + num_workers=8, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", +) +sampler_config = SamplerConfig( + dataloader_cfg=dataloader_cfg, + prompt_repeat_k=prompt_repeat_k, +) +training_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, +) +agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=training_sample_params, +) +produce_strategy_config = SyncProduceStrategyConfig() +agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="train_task", + agent_loop_config=agent_loop_config, + produce_strategy_config=produce_strategy_config, + sampler_config=sampler_config, + ), +) + +# 5. evaluation +eval_agent_loop_manager_cfg = None +evaluator_config = None +if enable_evaluate: + with open(eval_meta_data_path, "r", encoding="utf-8") as f: + eval_ds_collections = json.load(f) + + eval_dataset_cfg = [] + for name, data in eval_ds_collections.items(): + annotations = _as_list(data["annotation"]) + for annotation in annotations: + eval_dataset_cfg.append( + { + "dataset": DatasetConfig( + name=name, + anno_path=annotation, + media_root=data.get("media_root", ""), + sample_ratio=data.get("sample_ratio", 1.0), + class_name="VLMJsonlDataset", + ), + "tokenize_fn": RLQwen3VLTokenizeFnConfig( + processor_path=model_path, + max_length=max_prompt_length, + system_message=data.get("system_message", None), + chat_template="qwen3-vl", + add_generation_prompt=True, + enable_thinking=True, + ignore_multimodal_info=True, + ), + } + ) + + eval_judger_config = ComposedJudgerConfig( + branches={ + "openai/gsm8k": GSM8KJudgerConfig( + judger_name="openai/gsm8k", + cpu_resources=CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1), + ), + "hiyouga/geometry3k": GEO3KJudgerConfig( + cpu_resources=CPUResourcesConfig(num_workers=1, num_cpus_per_worker=1), + ), + } + ) + eval_dataloader_cfg = DataloaderConfig( + dataset_config_list=eval_dataset_cfg, + num_workers=8, + pack_max_length=pack_max_length, + collator="fake_collator", + pack_level="none", + ) + eval_sampler_config = SamplerConfig( + dataloader_cfg=eval_dataloader_cfg, + prompt_repeat_k=eval_prompt_repeat_k, + ) + evaluation_sample_params = SampleParams( + max_tokens=max_response_length, + top_k=0, + top_p=1.0, + temperature=1.0, + min_tokens=0, + skip_special_tokens=False, + return_routed_experts=False, + ) + eval_agent_loop_config = SingleTurnAgentLoopConfig( + hf_checkpoint=model_path, + sample_params=evaluation_sample_params, + ) + eval_agent_loop_manager_cfg = AgentLoopManagerConfig( + tasks=TaskSpecConfig( + task_name="eval_task", + agent_loop_config=eval_agent_loop_config, + judger_config=eval_judger_config, + sampler_config=eval_sampler_config, + ), + ) + evaluator_config = EvaluatorConfig() + +# 6. multi-teacher pure on-policy distillation +# +# This is the only topology block that needs to be edited for an experiment. +# Every training record's data_source must have an entry in +# data_source_teacher_map. Teacher model paths are read from +# GSM8K_TEACHER_MODEL_PATH and GEO3K_TEACHER_MODEL_PATH. +opd_config = OPDConfig( + mode="pg-opd", + task_adv_weight=0.0, + opd_adv_weight=1.0, + teachers=[ + OPDTeacherConfig( + name="gsm8k_teacher", + endpoint="http://127.0.0.1:13141", + launch_config=OPDTeacherLaunchConfig( + model_path=gsm8k_teacher_model_path, + cuda_visible_devices="6", + tensor_parallel_size=1, + expert_parallel_size=1, + context_length=max_num_tokens, + max_batch_size=max_num_tokens, + gpu_memory_utilization=0.8, + ), + ), + OPDTeacherConfig( + name="geo3k_teacher", + endpoint="http://127.0.0.1:13142", + launch_config=OPDTeacherLaunchConfig( + model_path=geo3k_teacher_model_path, + cuda_visible_devices="7", + tensor_parallel_size=1, + expert_parallel_size=1, + context_length=max_num_tokens, + max_batch_size=max_num_tokens, + gpu_memory_utilization=0.8, + ), + ), + ], + data_source_teacher_map={ + "openai/gsm8k": "gsm8k_teacher", + "hiyouga/geometry3k": "geo3k_teacher", + }, +) + +trainer = RLColocateTrainerConfig( + resources=resources, + train_worker_cfg=train_worker_cfg, # TODO: uniform naming of cfg and config + rollout_config=rollout_config, + tokenizer_path=model_path, + replay_buffer_config=SyncReplayBufferConfig(), + agent_loop_manager_cfg=agent_loop_manager_cfg, + eval_agent_loop_manager_cfg=eval_agent_loop_manager_cfg, + evaluator_config=evaluator_config, + load_from=model_path, + train_batch_size=train_batch_size, + advantage_estimator_config=GRPOAdvantageConfig(eps=1e-8), + opd_config=opd_config, + enable_evaluate=enable_evaluate, + enable_initial_evaluate=enable_evaluate, + evaluate_step=evaluate_step, + total_epochs=total_epochs, + work_dir=work_dir, + seed=1234, + debug_rollout=False, +) diff --git a/recipe/on_policy_distillation/rl_dapo_math_opd.py b/recipe/on_policy_distillation/config/rl_dapo_math_opd.py similarity index 92% rename from recipe/on_policy_distillation/rl_dapo_math_opd.py rename to recipe/on_policy_distillation/config/rl_dapo_math_opd.py index 70351650f3..a548ebcc62 100644 --- a/recipe/on_policy_distillation/rl_dapo_math_opd.py +++ b/recipe/on_policy_distillation/config/rl_dapo_math_opd.py @@ -2,7 +2,6 @@ from pathlib import Path from transformers import AutoTokenizer - from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig from xtuner.v1.data_proto.rl_data import SampleParams from xtuner.v1.datasets.config import DataloaderConfig, DatasetConfig @@ -19,16 +18,25 @@ from xtuner.v1.rl.evaluator import EvaluatorConfig from xtuner.v1.rl.judger import DapoMathJudgerConfig from xtuner.v1.rl.loss import GRPOLossConfig -from xtuner.v1.rl.on_policy_distillation import OPDConfig, OPDTeacherConfig +from xtuner.v1.rl.on_policy_distillation import ( + OPDConfig, + OPDTeacherConfig, + OPDTeacherLaunchConfig, +) from xtuner.v1.rl.replay_buffer import SyncReplayBufferConfig from xtuner.v1.rl.rollout.worker import RolloutConfig from xtuner.v1.rl.trainer import WorkerConfig -from xtuner.v1.rl.utils import AcceleratorResourcesConfig, CPUResourcesConfig, get_eos_token +from xtuner.v1.rl.utils import ( + AcceleratorResourcesConfig, + CPUResourcesConfig, + get_eos_token, +) from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig work_dir = os.environ["WORK_DIR"] model_path = os.environ["MODEL_PATH"] +teacher_model_path = os.environ["TEACHER_MODEL_PATH"] data_path = os.environ["DATA_PATH"] eval_data_path = os.environ.get("EVAL_DATA_PATH", "") NNODE = int(os.environ.get("WORLD_SIZE", "1")) @@ -202,7 +210,16 @@ mode="pg-opd", task_adv_weight=0.0, opd_adv_weight=1.0, - teachers=[OPDTeacherConfig(name="teacher", endpoint="http://127.0.0.1:13141")], + teachers=[ + OPDTeacherConfig( + name="teacher", + endpoint="http://127.0.0.1:13141", + launch_config=OPDTeacherLaunchConfig( + model_path=teacher_model_path, + cuda_visible_devices="7", + ), + ) + ], data_source_teacher_map={"math_dapo": "teacher"}, ) diff --git a/recipe/on_policy_distillation/run_sampled_token_opd.sh b/recipe/on_policy_distillation/run_sampled_token_opd.sh deleted file mode 100755 index a2a5e28e0b..0000000000 --- a/recipe/on_policy_distillation/run_sampled_token_opd.sh +++ /dev/null @@ -1,136 +0,0 @@ -#!/usr/bin/env bash - -set -euo pipefail - -SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) -REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../.." && pwd) - -STUDENT_MODEL_PATH="$1" -TEACHER_MODEL_PATH="$2" -DATA_PATH="$3" - -STUDENT_CUDA_VISIBLE_DEVICES="0,1,2,3" -TEACHER_CUDA_VISIBLE_DEVICES="7" -TEACHER_HOST="127.0.0.1" -TEACHER_PORT="13141" -TEACHER_TP_SIZE="1" -TEACHER_CHUNKED_PREFILL_SIZE="4096" -TEACHER_GPU_MEMORY_UTILIZATION="0.6" -TEACHER_STARTUP_TIMEOUT_S="1200" - -USE_SGLANG=${XTUNER_USE_SGLANG:-0} -USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} -USE_VLLM=${XTUNER_USE_VLLM:-0} - -if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then - OPD_BACKEND="sglang" - TEACHER_HEALTH_PATH="health_generate" - TEACHER_MODEL_INFO_PATH="get_model_info" -elif [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then - OPD_BACKEND="lmdeploy" - TEACHER_HEALTH_PATH="health" - TEACHER_MODEL_INFO_PATH="v1/models" -else - echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 - exit 1 -fi - -export XTUNER_USE_SGLANG="${USE_SGLANG}" -export XTUNER_USE_LMDEPLOY="${USE_LMDEPLOY}" -export XTUNER_USE_VLLM="${USE_VLLM}" - -WORK_DIR="${REPO_ROOT}/work_dirs/dapo_math_opd" -TEACHER_ENDPOINT="http://${TEACHER_HOST}:${TEACHER_PORT}" -TEACHER_LOG_FILE="${WORK_DIR}/teacher.log" - -mkdir -p "${WORK_DIR}" - -TEACHER_PID="" -TRAINING_STARTED=0 - -cleanup() { - local exit_code=$? - local attempt - - trap - EXIT INT TERM - - if [[ -n "${TEACHER_PID}" ]]; then - kill -TERM -- "-${TEACHER_PID}" 2>/dev/null || true - for ((attempt = 0; attempt < 30; attempt++)); do - if ! kill -0 -- "-${TEACHER_PID}" 2>/dev/null; then - break - fi - sleep 1 - done - if kill -0 -- "-${TEACHER_PID}" 2>/dev/null; then - kill -KILL -- "-${TEACHER_PID}" 2>/dev/null || true - fi - wait "${TEACHER_PID}" 2>/dev/null || true - fi - - if (( TRAINING_STARTED )); then - ray stop --force >/dev/null 2>&1 || true - fi - - exit "${exit_code}" -} - -wait_for_teacher() { - local deadline=$((SECONDS + TEACHER_STARTUP_TIMEOUT_S)) - - while (( SECONDS < deadline )); do - if ! kill -0 "${TEACHER_PID}" 2>/dev/null; then - echo "Teacher process exited before becoming ready." >&2 - tail -n 50 "${TEACHER_LOG_FILE}" >&2 || true - return 1 - fi - if curl -sf --max-time 2 "${TEACHER_ENDPOINT}/${TEACHER_HEALTH_PATH}" >/dev/null; then - return 0 - fi - echo "Waiting for teacher service at ${TEACHER_ENDPOINT}..." - sleep 5 - done - - echo "Teacher service did not become ready within ${TEACHER_STARTUP_TIMEOUT_S} seconds." >&2 - tail -n 50 "${TEACHER_LOG_FILE}" >&2 || true - return 1 -} - -trap cleanup EXIT -trap "exit 130" INT -trap "exit 143" TERM - -echo "Starting teacher model: ${TEACHER_MODEL_PATH}" -echo "Teacher backend: ${OPD_BACKEND}" -echo "Teacher GPUs: ${TEACHER_CUDA_VISIBLE_DEVICES}" -echo "Teacher log: ${TEACHER_LOG_FILE}" - -setsid env \ - CUDA_VISIBLE_DEVICES="${TEACHER_CUDA_VISIBLE_DEVICES}" \ - PYTHONUNBUFFERED=1 \ - TEACHER_HOST="${TEACHER_HOST}" \ - TEACHER_PORT="${TEACHER_PORT}" \ - TEACHER_TP_SIZE="${TEACHER_TP_SIZE}" \ - TEACHER_CHUNKED_PREFILL_SIZE="${TEACHER_CHUNKED_PREFILL_SIZE}" \ - TEACHER_GPU_MEMORY_UTILIZATION="${TEACHER_GPU_MEMORY_UTILIZATION}" \ - bash "${SCRIPT_DIR}/start_pg_opd_teacher.sh" "${TEACHER_MODEL_PATH}" \ - >"${TEACHER_LOG_FILE}" 2>&1 & -TEACHER_PID=$! - -wait_for_teacher -curl -sS --max-time 10 "${TEACHER_ENDPOINT}/${TEACHER_MODEL_INFO_PATH}" -echo -echo "Teacher service is ready at ${TEACHER_ENDPOINT}" -echo "Starting Pure PG-OPD training with student GPUs: ${STUDENT_CUDA_VISIBLE_DEVICES}" - -export WORK_DIR -export PYTHONUNBUFFERED=1 - -cd "${REPO_ROOT}" -TRAINING_STARTED=1 -CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES}" \ - bash -o pipefail examples/v1/scripts/run_rl.sh \ - recipe/on_policy_distillation/rl_dapo_math_opd.py \ - "${OPD_BACKEND}" \ - "${STUDENT_MODEL_PATH}" \ - "${DATA_PATH}" diff --git a/recipe/on_policy_distillation/scripts/launch_teacher_utils.sh b/recipe/on_policy_distillation/scripts/launch_teacher_utils.sh new file mode 100644 index 0000000000..1a6134f833 --- /dev/null +++ b/recipe/on_policy_distillation/scripts/launch_teacher_utils.sh @@ -0,0 +1,247 @@ +#!/usr/bin/env bash + +# Source-only helpers for launching, waiting for, and stopping OPD Teacher servers. + +OPD_RECIPE_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd) + +TEACHER_NAMES=() +TEACHER_ENDPOINTS=() +TEACHER_HEALTH_URLS=() +TEACHER_MODEL_INFO_URLS=() +TEACHER_LOG_FILES=() +TEACHER_PIDS=() + +start_single_teacher_server() { + _start_teacher_servers "$1" "$2" "$3" "1" "1" +} + +start_teacher_servers() { + _start_teacher_servers "$1" "$2" "$3" "" "0" +} + +wait_for_teacher_servers() { + local startup_timeout_s=$1 + local deadline=$((SECONDS + startup_timeout_s)) + local teacher_index + local health_check_index + local pid + local log_file + local -a teacher_ready=() + local -a health_check_pids=() + local -a health_check_indices=() + local -a pending_names=() + + if (( ${#TEACHER_NAMES[@]} == 0 )); then + echo "No Teacher servers have been configured." >&2 + return 1 + fi + + for teacher_index in "${!TEACHER_NAMES[@]}"; do + teacher_ready[teacher_index]=0 + done + + while (( SECONDS < deadline )); do + for teacher_index in "${!TEACHER_NAMES[@]}"; do + pid="${TEACHER_PIDS[teacher_index]}" + log_file="${TEACHER_LOG_FILES[teacher_index]}" + if [[ -n "${pid}" ]] && ! kill -0 "${pid}" 2>/dev/null; then + echo "Teacher ${TEACHER_NAMES[teacher_index]} exited before becoming ready." >&2 + if [[ -n "${log_file}" ]]; then + tail -n 50 "${log_file}" >&2 || true + fi + return 1 + fi + done + + health_check_pids=() + health_check_indices=() + for teacher_index in "${!TEACHER_NAMES[@]}"; do + if (( teacher_ready[teacher_index] )); then + continue + fi + curl -sf --max-time 2 \ + "${TEACHER_HEALTH_URLS[teacher_index]}" \ + >/dev/null 2>&1 & + health_check_pids+=("$!") + health_check_indices+=("${teacher_index}") + done + + for health_check_index in "${!health_check_pids[@]}"; do + teacher_index=${health_check_indices[health_check_index]} + if wait "${health_check_pids[health_check_index]}"; then + teacher_ready[teacher_index]=1 + echo "Teacher ${TEACHER_NAMES[teacher_index]} is ready at ${TEACHER_ENDPOINTS[teacher_index]}" + fi + done + + pending_names=() + for teacher_index in "${!TEACHER_NAMES[@]}"; do + if (( ! teacher_ready[teacher_index] )); then + pending_names+=("${TEACHER_NAMES[teacher_index]}") + fi + done + if (( ${#pending_names[@]} == 0 )); then + return 0 + fi + + echo "Waiting for teachers: ${pending_names[*]}" + sleep 5 + done + + for teacher_index in "${!TEACHER_NAMES[@]}"; do + if (( teacher_ready[teacher_index] )); then + continue + fi + echo "Teacher ${TEACHER_NAMES[teacher_index]} did not become ready within ${startup_timeout_s}s." >&2 + log_file="${TEACHER_LOG_FILES[teacher_index]}" + if [[ -n "${log_file}" ]]; then + tail -n 50 "${log_file}" >&2 || true + fi + done + return 1 +} + +stop_teacher_servers() { + local attempt + local any_alive + local pid + + for pid in "${TEACHER_PIDS[@]}"; do + if [[ -n "${pid}" ]]; then + kill -TERM -- "-${pid}" 2>/dev/null || true + fi + done + + for ((attempt = 0; attempt < 30; attempt++)); do + any_alive=0 + for pid in "${TEACHER_PIDS[@]}"; do + if [[ -n "${pid}" ]] && kill -0 -- "-${pid}" 2>/dev/null; then + any_alive=1 + break + fi + done + if (( ! any_alive )); then + break + fi + sleep 1 + done + + for pid in "${TEACHER_PIDS[@]}"; do + if [[ -n "${pid}" ]]; then + if kill -0 -- "-${pid}" 2>/dev/null; then + kill -KILL -- "-${pid}" 2>/dev/null || true + fi + wait "${pid}" 2>/dev/null || true + fi + done + + TEACHER_PIDS=() +} + +_start_teacher_servers() { + local config_file=$1 + local backend=$2 + local work_dir=$3 + local expected_teacher_count=$4 + local require_local_teacher=$5 + local teacher_count + local teacher_index + local teacher_field_offset=1 + local teacher_name + local teacher_safe_name + local teacher_endpoint + local teacher_health_url + local teacher_model_info_url + local teacher_command_arg_count + local teacher_command_offset + local teacher_log_file + local -a teacher_fields=() + local -a teacher_command=() + + _reset_teacher_server_state + mkdir -p "${work_dir}" + + # Command builder output: + # count, then repeated name, safe name, endpoint, health URL, + # model-info URL, command-argument count, and command argv. + mapfile -d "" -t teacher_fields < <( + python "${OPD_RECIPE_DIR}/build_teacher_server_commands.py" \ + "${config_file}" "${backend}" + ) + if (( ${#teacher_fields[@]} == 0 )); then + echo "Teacher command builder did not return any records." >&2 + return 1 + fi + + teacher_count=${teacher_fields[0]} + if [[ -n "${expected_teacher_count}" ]] && (( teacher_count != expected_teacher_count )); then + echo "Expected ${expected_teacher_count} Teacher, got ${teacher_count}." >&2 + return 1 + fi + if (( teacher_count == 0 )); then + echo "OPD config does not contain any Teachers: ${config_file}" >&2 + return 1 + fi + + echo "Teacher backend: ${backend}" + for ((teacher_index = 0; teacher_index < teacher_count; teacher_index++)); do + teacher_name=${teacher_fields[teacher_field_offset]} + teacher_safe_name=${teacher_fields[teacher_field_offset + 1]} + teacher_endpoint=${teacher_fields[teacher_field_offset + 2]} + teacher_health_url=${teacher_fields[teacher_field_offset + 3]} + teacher_model_info_url=${teacher_fields[teacher_field_offset + 4]} + teacher_command_arg_count=${teacher_fields[teacher_field_offset + 5]} + teacher_command_offset=$((teacher_field_offset + 6)) + teacher_command=( + "${teacher_fields[@]:teacher_command_offset:teacher_command_arg_count}" + ) + teacher_field_offset=$((teacher_command_offset + teacher_command_arg_count)) + + TEACHER_NAMES+=("${teacher_name}") + TEACHER_ENDPOINTS+=("${teacher_endpoint}") + TEACHER_HEALTH_URLS+=("${teacher_health_url}") + TEACHER_MODEL_INFO_URLS+=("${teacher_model_info_url}") + + if (( teacher_command_arg_count == 0 )); then + if (( require_local_teacher )); then + echo "Teacher ${teacher_name} must define launch_config for local startup." >&2 + return 1 + fi + echo "Using externally managed Teacher ${teacher_name} at ${teacher_endpoint}" + TEACHER_LOG_FILES+=("") + TEACHER_PIDS+=("") + continue + fi + + if [[ -n "${expected_teacher_count}" ]]; then + teacher_log_file="${work_dir}/teacher.log" + else + teacher_log_file="${work_dir}/teacher_${teacher_index}_${teacher_safe_name}.log" + fi + TEACHER_LOG_FILES+=("${teacher_log_file}") + + echo "Starting Teacher ${teacher_name}" + echo "Teacher endpoint: ${teacher_endpoint}" + echo "Teacher log: ${teacher_log_file}" + + setsid env \ + PYTHONUNBUFFERED=1 \ + "${teacher_command[@]}" \ + >"${teacher_log_file}" 2>&1 & + TEACHER_PIDS+=("$!") + done + + if (( teacher_field_offset != ${#teacher_fields[@]} )); then + echo "Teacher command builder returned malformed records." >&2 + return 1 + fi +} + +_reset_teacher_server_state() { + TEACHER_NAMES=() + TEACHER_ENDPOINTS=() + TEACHER_HEALTH_URLS=() + TEACHER_MODEL_INFO_URLS=() + TEACHER_LOG_FILES=() + TEACHER_PIDS=() +} diff --git a/recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh b/recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh new file mode 100755 index 0000000000..4241b4d1a2 --- /dev/null +++ b/recipe/on_policy_distillation/scripts/run_sampled_token_mopd.sh @@ -0,0 +1,75 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../../.." && pwd) +source "${SCRIPT_DIR}/launch_teacher_utils.sh" + +export STUDENT_MODEL_PATH=${1:?"Usage: $0 STUDENT_MODEL_PATH DATA_PATH"} +export MODEL_PATH="${STUDENT_MODEL_PATH}" +export DATA_PATH=${2:?"Usage: $0 STUDENT_MODEL_PATH DATA_PATH"} +export GSM8K_TEACHER_MODEL_PATH=${GSM8K_TEACHER_MODEL_PATH:?"GSM8K_TEACHER_MODEL_PATH is required"} +export GEO3K_TEACHER_MODEL_PATH=${GEO3K_TEACHER_MODEL_PATH:?"GEO3K_TEACHER_MODEL_PATH is required"} + +export OPD_CONFIG_PATH="${OPD_CONFIG_PATH:-recipe/on_policy_distillation/config/rl_dapo_math_mopd.py}" +export EVAL_DATA_PATH="${EVAL_DATA_PATH:-}" +export STUDENT_CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES:-0,1,2,3}" +export TEACHER_STARTUP_TIMEOUT_S="${TEACHER_STARTUP_TIMEOUT_S:-1200}" +export PYTHONUNBUFFERED=1 + +USE_SGLANG=${XTUNER_USE_SGLANG:-0} +USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} +USE_VLLM=${XTUNER_USE_VLLM:-0} + +if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="sglang" +elif [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="lmdeploy" +else + echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 + exit 1 +fi + +export XTUNER_USE_SGLANG="${USE_SGLANG}" +export XTUNER_USE_LMDEPLOY="${USE_LMDEPLOY}" +export XTUNER_USE_VLLM="${USE_VLLM}" + +export WORK_DIR="${WORK_DIR:-${REPO_ROOT}/work_dirs/dapo_math_mopd}" +export OPD_CONFIG_FILE="${REPO_ROOT}/${OPD_CONFIG_PATH}" + +TRAINING_STARTED=0 + +cleanup() { + local exit_code=$? + + trap - EXIT INT TERM + + stop_teacher_servers + + if (( TRAINING_STARTED )); then + ray stop --force >/dev/null 2>&1 || true + fi + + exit "${exit_code}" +} + +trap cleanup EXIT +trap "exit 130" INT +trap "exit 143" TERM + +start_teacher_servers "${OPD_CONFIG_FILE}" "${OPD_BACKEND}" "${WORK_DIR}" +wait_for_teacher_servers "${TEACHER_STARTUP_TIMEOUT_S}" + +echo "All ${#TEACHER_NAMES[@]} teachers are ready." +echo "Starting MOPD training with student GPUs: ${STUDENT_CUDA_VISIBLE_DEVICES}" + +cd "${REPO_ROOT}" +TRAINING_STARTED=1 +CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES}" \ + bash -o pipefail examples/v1/scripts/run_rl.sh \ + "${OPD_CONFIG_FILE}" \ + "${OPD_BACKEND}" \ + "${STUDENT_MODEL_PATH}" \ + "${DATA_PATH}" \ + "${EVAL_DATA_PATH}" diff --git a/recipe/on_policy_distillation/scripts/run_sampled_token_opd.sh b/recipe/on_policy_distillation/scripts/run_sampled_token_opd.sh new file mode 100755 index 0000000000..cff10eceea --- /dev/null +++ b/recipe/on_policy_distillation/scripts/run_sampled_token_opd.sh @@ -0,0 +1,72 @@ +#!/usr/bin/env bash + +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../../.." && pwd) +source "${SCRIPT_DIR}/launch_teacher_utils.sh" + +export STUDENT_MODEL_PATH=${1:?"Usage: $0 STUDENT_MODEL_PATH TEACHER_MODEL_PATH DATA_PATH"} +export TEACHER_MODEL_PATH=${2:?"Usage: $0 STUDENT_MODEL_PATH TEACHER_MODEL_PATH DATA_PATH"} +export DATA_PATH=${3:?"Usage: $0 STUDENT_MODEL_PATH TEACHER_MODEL_PATH DATA_PATH"} +export OPD_CONFIG_PATH="${OPD_CONFIG_PATH:-recipe/on_policy_distillation/config/rl_dapo_math_opd.py}" +export EVAL_DATA_PATH="${EVAL_DATA_PATH:-}" +export STUDENT_CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES:-0,1,2,3}" +export TEACHER_STARTUP_TIMEOUT_S="${TEACHER_STARTUP_TIMEOUT_S:-1200}" +export PYTHONUNBUFFERED=1 + +USE_SGLANG=${XTUNER_USE_SGLANG:-0} +USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} +USE_VLLM=${XTUNER_USE_VLLM:-0} + +if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="sglang" +elif [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then + OPD_BACKEND="lmdeploy" +else + echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 + exit 1 +fi + +export XTUNER_USE_SGLANG="${USE_SGLANG}" +export XTUNER_USE_LMDEPLOY="${USE_LMDEPLOY}" +export XTUNER_USE_VLLM="${USE_VLLM}" + +export WORK_DIR="${WORK_DIR:-${REPO_ROOT}/work_dirs/dapo_math_opd}" +export OPD_CONFIG_FILE="${REPO_ROOT}/${OPD_CONFIG_PATH}" + +TRAINING_STARTED=0 + +cleanup() { + local exit_code=$? + + trap - EXIT INT TERM + + stop_teacher_servers + + if (( TRAINING_STARTED )); then + ray stop --force >/dev/null 2>&1 || true + fi + + exit "${exit_code}" +} + +trap cleanup EXIT +trap "exit 130" INT +trap "exit 143" TERM + +start_single_teacher_server "${OPD_CONFIG_FILE}" "${OPD_BACKEND}" "${WORK_DIR}" +wait_for_teacher_servers "${TEACHER_STARTUP_TIMEOUT_S}" +curl -sS --max-time 10 "${TEACHER_MODEL_INFO_URLS[0]}" +echo +echo "Starting Pure PG-OPD training with student GPUs: ${STUDENT_CUDA_VISIBLE_DEVICES}" + +cd "${REPO_ROOT}" +TRAINING_STARTED=1 +CUDA_VISIBLE_DEVICES="${STUDENT_CUDA_VISIBLE_DEVICES}" \ + bash -o pipefail examples/v1/scripts/run_rl.sh \ + "${OPD_CONFIG_FILE}" \ + "${OPD_BACKEND}" \ + "${STUDENT_MODEL_PATH}" \ + "${DATA_PATH}" \ + "${EVAL_DATA_PATH}" diff --git a/recipe/on_policy_distillation/start_pg_opd_teacher.sh b/recipe/on_policy_distillation/start_pg_opd_teacher.sh deleted file mode 100755 index 477bffcc14..0000000000 --- a/recipe/on_policy_distillation/start_pg_opd_teacher.sh +++ /dev/null @@ -1,45 +0,0 @@ -#!/usr/bin/env bash - -set -euo pipefail - -SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) -REPO_ROOT=$(cd -- "${SCRIPT_DIR}/../.." && pwd) - -TEACHER_MODEL_PATH=${1:?"Usage: $0 TEACHER_MODEL_PATH"} -TEACHER_HOST=${TEACHER_HOST:-127.0.0.1} -TEACHER_PORT=${TEACHER_PORT:-13141} -TEACHER_TP_SIZE=${TEACHER_TP_SIZE:-1} -TEACHER_CHUNKED_PREFILL_SIZE=${TEACHER_CHUNKED_PREFILL_SIZE:-4096} -TEACHER_GPU_MEMORY_UTILIZATION=${TEACHER_GPU_MEMORY_UTILIZATION:-0.6} -PYTHON_EXECUTABLE=${PYTHON_EXECUTABLE:-python} -LMDEPLOY_PATH=${LMDEPLOY_PATH:-"${REPO_ROOT}/work_dirs/lmdeploy"} - -USE_SGLANG=${XTUNER_USE_SGLANG:-0} -USE_LMDEPLOY=${XTUNER_USE_LMDEPLOY:-0} -USE_VLLM=${XTUNER_USE_VLLM:-0} - -if [[ "${USE_SGLANG}" == "1" && "${USE_LMDEPLOY}" == "0" && "${USE_VLLM}" == "0" ]]; then - exec "${PYTHON_EXECUTABLE}" -m sglang.launch_server \ - --model-path "${TEACHER_MODEL_PATH}" \ - --host "${TEACHER_HOST}" \ - --port "${TEACHER_PORT}" \ - --tp "${TEACHER_TP_SIZE}" \ - --chunked-prefill-size "${TEACHER_CHUNKED_PREFILL_SIZE}" \ - --mem-fraction-static "${TEACHER_GPU_MEMORY_UTILIZATION}" -fi - -if [[ "${USE_SGLANG}" == "0" && "${USE_LMDEPLOY}" == "1" && "${USE_VLLM}" == "0" ]]; then - export PYTHONPATH="${LMDEPLOY_PATH}${PYTHONPATH:+:${PYTHONPATH}}" - exec "${PYTHON_EXECUTABLE}" -m lmdeploy serve api_server "${TEACHER_MODEL_PATH}" \ - --backend pytorch \ - --role Hybrid \ - --logprobs-mode raw_logprobs \ - --server-name "${TEACHER_HOST}" \ - --server-port "${TEACHER_PORT}" \ - --tp "${TEACHER_TP_SIZE}" \ - --max-prefill-token-num "${TEACHER_CHUNKED_PREFILL_SIZE}" \ - --cache-max-entry-count "${TEACHER_GPU_MEMORY_UTILIZATION}" -fi - -echo "Exactly one of XTUNER_USE_SGLANG and XTUNER_USE_LMDEPLOY must be set to 1; XTUNER_USE_VLLM must be 0." >&2 -exit 1 diff --git a/tests/rl/test_on_policy_distillation.py b/tests/rl/test_on_policy_distillation.py index 5f23c22069..fa627bbee4 100644 --- a/tests/rl/test_on_policy_distillation.py +++ b/tests/rl/test_on_policy_distillation.py @@ -10,23 +10,29 @@ import httpx import torch +from recipe.on_policy_distillation.build_teacher_server_commands import ( + build_teacher_server_command, +) from xtuner.v1.data_proto.rl_data import RolloutState, Status from xtuner.v1.rl.loss import GRPOLossConfig from xtuner.v1.rl.on_policy_distillation import ( OPDConfig, OPDTeacherConfig, + OPDTeacherLaunchConfig, TeacherLogprobClient, apply_opd_kl_to_advantages, ) from xtuner.v1.rl.utils import find_free_ports + REPO_ROOT = Path(__file__).resolve().parents[2] BASELINE_PATH = os.getenv("XTUNER_OPD_BASELINE") TEACHER_MODEL_PATH = os.getenv("XTUNER_OPD_TEACHER_MODEL") STUDENT_MODEL_PATH = os.getenv("XTUNER_OPD_STUDENT_MODEL") TEACHER_STARTUP_TIMEOUT_S = float(os.getenv("XTUNER_OPD_TEACHER_STARTUP_TIMEOUT_S", "1200")) -TEACHER_START_SCRIPT = REPO_ROOT / "recipe/on_policy_distillation/start_pg_opd_teacher.sh" -TRAINER_CONFIG_PATH = REPO_ROOT / "recipe/on_policy_distillation/rl_dapo_math_opd.py" +TRAINER_CONFIG_PATH = ( + REPO_ROOT / "recipe/on_policy_distillation/config/rl_dapo_math_opd.py" +) def _wait_for_teacher(process: subprocess.Popen, endpoint: str, backend: str) -> None: @@ -286,6 +292,7 @@ def build_capturing_training_workers(self, placement_group): trainer_environment = { "WORK_DIR": str(run_dir / "trainer"), "MODEL_PATH": str(student_model_path), + "TEACHER_MODEL_PATH": os.getenv("XTUNER_OPD_TEACHER_MODEL", str(student_model_path)), "DATA_PATH": os.getenv("XTUNER_OPD_DATA_PATH", str(baseline_path)), "WORLD_SIZE": "1", "ONLY_CALC_MISMATCH_RATIO": "1", @@ -487,10 +494,22 @@ def setUpClass(cls) -> None: cls.teacher_endpoint = f"http://127.0.0.1:{port}" cls.teacher_backend = TeacherLogprobClient._resolve_backend_from_env() teacher_env = os.environ.copy() - teacher_env["TEACHER_HOST"] = "127.0.0.1" - teacher_env["TEACHER_PORT"] = str(port) + teacher_command = build_teacher_server_command( + OPDTeacherConfig( + name="teacher", + endpoint=cls.teacher_endpoint, + launch_config=OPDTeacherLaunchConfig( + model_path=str(TEACHER_MODEL_PATH), + cuda_visible_devices=teacher_env.get("CUDA_VISIBLE_DEVICES") + or "7", + ), + ), + cls.teacher_backend, + ) + if not teacher_command: + raise ValueError("Teacher must define launch_config for local startup") cls.teacher_process = subprocess.Popen( - ["bash", str(TEACHER_START_SCRIPT), str(TEACHER_MODEL_PATH)], + teacher_command, cwd=REPO_ROOT, env=teacher_env, start_new_session=True, diff --git a/xtuner/v1/rl/on_policy_distillation.py b/xtuner/v1/rl/on_policy_distillation.py index 9ca18a2f30..66a9dece16 100644 --- a/xtuner/v1/rl/on_policy_distillation.py +++ b/xtuner/v1/rl/on_policy_distillation.py @@ -4,6 +4,7 @@ import math import os import time +from pathlib import Path from typing import Any, Literal, cast import httpx @@ -14,6 +15,27 @@ from xtuner.v1.rl.loss.base_loss import BaseRLLossContext +class OPDTeacherLaunchConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + model_path: str | Path + cuda_visible_devices: str + dtype: Literal["auto", "float16", "bfloat16"] = "bfloat16" + tensor_parallel_size: int = Field(default=1, gt=0) + expert_parallel_size: int = Field(default=1, gt=0) + context_length: int | None = Field(default=None, gt=0) + max_batch_size: int | None = Field(default=None, gt=0) + chunked_prefill_size: int | None = Field(default=4096, gt=0) + max_prefill_token_num: int | None = Field(default=4096, gt=0) + gpu_memory_utilization: float = Field(default=0.6, gt=0.0, le=1.0) + + @model_validator(mode="after") + def validate_parallel_sizes(self) -> OPDTeacherLaunchConfig: + if self.tensor_parallel_size > 1 and self.expert_parallel_size > 1: + raise ValueError("tensor_parallel_size and expert_parallel_size cannot both be greater than 1") + return self + + class OPDTeacherConfig(BaseModel): model_config = ConfigDict(extra="forbid") @@ -23,6 +45,7 @@ class OPDTeacherConfig(BaseModel): request_timeout_s: float = Field(default=1200.0, gt=0.0) max_retry_per_sample: int = Field(default=2, ge=0) max_concurrency: int = Field(default=128, gt=0) + launch_config: OPDTeacherLaunchConfig | None = None class OPDConfig(BaseModel):