From 6f1c9fc77b0215732cdbcaa9761a6327f8ac1ef7 Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Fri, 29 May 2026 20:46:10 -0700 Subject: [PATCH] [RL] Fix crash when the reqs in a batch have a mix of `return_routed_experts` = True and False. (#26423) Co-authored-by: root Co-authored-by: Cursor Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com> --- .../scheduler_components/output_streamer.py | 53 ++++++-- .../sglang/srt/managers/tokenizer_manager.py | 4 +- .../rl/test_return_routed_experts.py | 114 +++++++++++++++--- 3 files changed, 141 insertions(+), 30 deletions(-) diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index 90fdebb1f..4bbde9887 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -121,8 +121,21 @@ class SchedulerOutputStreamer: skip_req: Optional[Req] = None, is_idle_batch: bool = False, ): + return_hidden_states = any( + req.return_hidden_states for req in reqs if req is not skip_req + ) + return_routed_experts = any( + req.return_routed_experts for req in reqs if req is not skip_req + ) + return_indexer_topk = any( + req.return_indexer_topk for req in reqs if req is not skip_req + ) + acc = _GenerationStreamAccumulator( return_logprob=return_logprob, + return_hidden_states=return_hidden_states, + return_routed_experts=return_routed_experts, + return_indexer_topk=return_indexer_topk, spec_algorithm=self.spec_algorithm, disaggregation_mode=self.disaggregation_mode, default_stream_interval=self.server_args.stream_interval, @@ -229,6 +242,9 @@ class SchedulerOutputStreamer: @dataclass(slots=True, kw_only=True) class _GenerationStreamAccumulator: return_logprob: bool + return_hidden_states: bool + return_routed_experts: bool + return_indexer_topk: bool spec_algorithm: Any disaggregation_mode: DisaggregationMode default_stream_interval: int @@ -256,9 +272,9 @@ class _GenerationStreamAccumulator: spec_num_correct_drafts: list = field(default_factory=list) spec_correct_drafts_histogram: list = field(default_factory=list) retraction_counts: list = field(default_factory=list) - output_hidden_states: list = field(default_factory=list) - routed_experts: list = field(default_factory=list) - indexer_topk: list = field(default_factory=list) + output_hidden_states: Optional[list] = None + routed_experts: Optional[list] = None + indexer_topk: Optional[list] = None customized_info: dict = field(default_factory=dict) time_stats: list = field(default_factory=list) input_token_logprobs_val: Optional[list] = None @@ -275,6 +291,13 @@ class _GenerationStreamAccumulator: output_token_ids_logprobs_idx: Optional[list] = None def __post_init__(self) -> None: + if self.return_hidden_states: + self.output_hidden_states = [] + if self.return_routed_experts: + self.routed_experts = [] + if self.return_indexer_topk: + self.indexer_topk = [] + if self.return_logprob: self.input_token_logprobs_val = [] self.input_token_logprobs_idx = [] @@ -434,12 +457,18 @@ class _GenerationStreamAccumulator: self.output_token_ids_logprobs_val.append([]) self.output_token_ids_logprobs_idx.append([]) - if req.return_hidden_states: - self.output_hidden_states.append(req.hidden_states) - if req.return_routed_experts: - self.routed_experts.append(req.routed_experts) - if req.return_indexer_topk: - self.indexer_topk.append(req.indexer_topk) + if self.return_hidden_states: + self.output_hidden_states.append( + req.hidden_states if req.return_hidden_states else None + ) + if self.return_routed_experts: + self.routed_experts.append( + req.routed_experts if req.return_routed_experts else None + ) + if self.return_indexer_topk: + self.indexer_topk.append( + req.indexer_topk if req.return_indexer_topk else None + ) if req.customized_info is not None: for k, v in req.customized_info.items(): @@ -486,9 +515,9 @@ class _GenerationStreamAccumulator: output_token_ids_logprobs_val=self.output_token_ids_logprobs_val, output_token_ids_logprobs_idx=self.output_token_ids_logprobs_idx, output_token_entropy_val=None, - output_hidden_states=self.output_hidden_states or None, - routed_experts=self.routed_experts or None, - indexer_topk=self.indexer_topk or None, + output_hidden_states=self.output_hidden_states, + routed_experts=self.routed_experts, + indexer_topk=self.indexer_topk, customized_info=self.customized_info, placeholder_tokens_idx=None, placeholder_tokens_val=None, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 2376811a8..656bfc7b6 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1784,7 +1784,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ] if getattr(recv_obj, "output_hidden_states", None): - meta_info["hidden_states"] = recv_obj.output_hidden_states[i] + hidden_states = recv_obj.output_hidden_states[i] + if hidden_states is not None: + meta_info["hidden_states"] = hidden_states if getattr(recv_obj, "routed_experts", None): val = recv_obj.routed_experts[i] if val is not None: diff --git a/test/registered/rl/test_return_routed_experts.py b/test/registered/rl/test_return_routed_experts.py index 5157616ef..90007074d 100644 --- a/test/registered/rl/test_return_routed_experts.py +++ b/test/registered/rl/test_return_routed_experts.py @@ -77,6 +77,32 @@ class TestReturnRoutedExperts(CustomTestCase): ] cls.reference_args = common cls.sampling_args = {"temperature": 0} + cls.texts = None + cls.baseline_results = None + cls.reference_results = None + cls._endpoints = [ + ( + "/generate", + cls._build_generate_payload, + extract_routed_experts_from_meta_info, + ), + ( + "/v1/chat/completions", + cls._build_chat_payload, + extract_routed_experts_from_openai_response, + ), + ( + "/v1/completions", + cls._build_completion_payload, + extract_routed_experts_from_openai_response, + ), + ] + + @classmethod + def _ensure_comparison_results(cls): + if cls.baseline_results is not None and cls.reference_results is not None: + return + # prepare ShareGPT dataset dataset_path = download_and_cache_hf_file(SHAREGPT_REPO_ID, SHAREGPT_FILENAME) with open(dataset_path) as f: @@ -96,23 +122,6 @@ class TestReturnRoutedExperts(CustomTestCase): if not cls.texts: raise ValueError("No valid texts found in the dataset") cls.texts = cls.texts[:100] - cls._endpoints = [ - ( - "/generate", - cls._build_generate_payload, - extract_routed_experts_from_meta_info, - ), - ( - "/v1/chat/completions", - cls._build_chat_payload, - extract_routed_experts_from_openai_response, - ), - ( - "/v1/completions", - cls._build_completion_payload, - extract_routed_experts_from_openai_response, - ), - ] cls.baseline_results = cls._collect_results(cls.baseline_args) cls.reference_results = cls._collect_results(cls.reference_args) @@ -128,8 +137,13 @@ class TestReturnRoutedExperts(CustomTestCase): def test_return_routed_experts_completions(cls): cls._run_endpoint_test("/v1/completions") + def test_mixed_return_routed_experts_batch_alignment(self): + self._run_mixed_batch_alignment_case([]) + self._run_mixed_batch_alignment_case(["--tokenizer-worker-num", 2]) + @classmethod def _run_endpoint_test(cls, endpoint): + cls._ensure_comparison_results() captured_baseline_experts = cls.baseline_results[endpoint] captured_reference_experts = cls.reference_results[endpoint] @@ -171,6 +185,72 @@ class TestReturnRoutedExperts(CustomTestCase): finally: kill_process_tree(process.pid) + @classmethod + def _run_mixed_batch_alignment_case(cls, other_args): + process = popen_launch_server( + DEFAULT_ENABLE_ROUTED_EXPERTS_MODEL_NAME_FOR_TEST, + DEFAULT_URL_FOR_TEST, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp", + 2, + "--enable-return-routed-experts", + "--disable-cuda-graph", + "--disable-piecewise-cuda-graph", + *other_args, + ], + ) + try: + responses = asyncio.run(cls._send_mixed_batch()) + cls._assert_mixed_batch_result(responses) + finally: + kill_process_tree(process.pid) + + @classmethod + async def _send_mixed_batch(cls): + payload_no_rr = { + "text": "The quick brown fox jumps over the lazy dog.", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 16, + "ignore_eos": True, + }, + "return_routed_experts": False, + } + payload_with_rr = { + "text": "The quick brown fox jumps over the lazy dog.", + "sampling_params": { + "temperature": 0, + "max_new_tokens": 16, + "ignore_eos": True, + }, + "return_routed_experts": True, + } + + async with aiohttp.ClientSession() as session: + return await asyncio.gather( + cls._post_generate(session, payload_no_rr), + cls._post_generate(session, payload_with_rr), + ) + + @staticmethod + async def _post_generate(session, payload): + async with session.post( + f"{DEFAULT_URL_FOR_TEST}/generate", json=payload + ) as response: + body = await response.json() + if response.status != 200: + raise AssertionError(f"HTTP {response.status}: {body}") + if "error" in body: + raise AssertionError(f"generate returned error: {body['error']}") + return body + + @classmethod + def _assert_mixed_batch_result(cls, responses): + no_rr, with_rr = responses + assert "routed_experts" not in no_rr.get("meta_info", {}) + assert "routed_experts" in with_rr.get("meta_info", {}) + @classmethod async def _collect_results_async(cls): results = {}