[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:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user