From a98d921658b2cb78ca1257a4d72a6a3969620456 Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Thu, 17 Sep 2026 09:21:09 -0700 Subject: [PATCH] [DP Attn] Fix crash for no token all-gather case (#39899) Co-authored-by: Byron Hsu --- python/sglang/srt/layers/logits_processor.py | 11 +++++--- .../model_executor/test_mlp_sync_pad_unpad.py | 26 +++++++++++++++++++ 2 files changed, 34 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 6aac6eb21..b488bfa9d 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -309,10 +309,15 @@ class LogitsMetadata: @classmethod def from_forward_batch(cls, forward_batch: ForwardBatch): + # MLP-sync may turn an idle rank into a dummy EXTEND for DP prefill + # graphs. It still has no real request whose last token needs logits. + forward_mode = forward_batch.forward_mode + if forward_batch._original_forward_mode == ForwardMode.IDLE: + forward_mode = ForwardMode.IDLE if ( - forward_batch.forward_mode.is_extend() + forward_mode.is_extend() and forward_batch.return_logprob - and not forward_batch.forward_mode.is_target_verify() + and not forward_mode.is_target_verify() ): extend_return_top_logprob = any( x > 0 for x in forward_batch.top_logprobs_nums @@ -340,7 +345,7 @@ class LogitsMetadata: draft_extend_select_index = None return cls( - forward_mode=forward_batch.forward_mode, + forward_mode=forward_mode, capture_hidden_mode=forward_batch.capture_hidden_mode, next_token_logits_buffer=forward_batch.next_token_logits_buffer, extend_return_logprob=extend_return_logprob, diff --git a/test/registered/unit/model_executor/test_mlp_sync_pad_unpad.py b/test/registered/unit/model_executor/test_mlp_sync_pad_unpad.py index 376982a86..a7c9565e1 100644 --- a/test/registered/unit/model_executor/test_mlp_sync_pad_unpad.py +++ b/test/registered/unit/model_executor/test_mlp_sync_pad_unpad.py @@ -15,6 +15,7 @@ from unittest.mock import MagicMock import torch +from sglang.srt.layers.logits_processor import LogitsMetadata, LogitsProcessor from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( DecodeCudaGraphRunner, @@ -40,6 +41,31 @@ def _logits_output(num_rows: int) -> SimpleNamespace: class TestMlpSyncPadUnpad(CustomTestCase): + def test_idle_rank_does_not_index_dummy_last_token(self): + # MLP-sync turns an idle rank into a dummy zero-token EXTEND batch. + empty = torch.empty(0, dtype=torch.int64) + batch = ForwardBatch( + forward_mode=ForwardMode.EXTEND, + batch_size=1, + input_ids=empty, + req_pool_indices=torch.tensor([0]), + seq_lens=torch.tensor([0]), + out_cache_loc=empty, + seq_lens_sum=0, + positions=empty, + extend_seq_lens=torch.tensor([0]), + extend_seq_lens_cpu=[0], + _original_forward_mode=ForwardMode.IDLE, + _original_batch_size=0, + ) + hidden = torch.empty(0, 4) + pruned, *_ = LogitsProcessor._get_pruned_states( + None, hidden, None, None, LogitsMetadata.from_forward_batch(batch) + ) + self.assertEqual(pruned.shape, (0, 4)) + # Attention and MLP execution still use the padded mode. + self.assertEqual(batch.forward_mode, ForwardMode.EXTEND) + def test_init_mlp_sync_metadata_scales_speculative_request_width(self): spec_info = SimpleNamespace( num_tokens_per_req=4,