[DP Attn] Fix crash for no token all-gather case (#39899)

Co-authored-by: Byron Hsu <byron+per@periodiclabs.ai>
This commit is contained in:
Byron Hsu
2026-09-17 09:21:09 -07:00
committed by GitHub
co-authored by Byron Hsu
parent 25c9f724d4
commit a98d921658
2 changed files with 34 additions and 3 deletions
+8 -3
View File
@@ -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,
@@ -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,