[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
|
@classmethod
|
||||||
def from_forward_batch(cls, forward_batch: ForwardBatch):
|
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 (
|
if (
|
||||||
forward_batch.forward_mode.is_extend()
|
forward_mode.is_extend()
|
||||||
and forward_batch.return_logprob
|
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(
|
extend_return_top_logprob = any(
|
||||||
x > 0 for x in forward_batch.top_logprobs_nums
|
x > 0 for x in forward_batch.top_logprobs_nums
|
||||||
@@ -340,7 +345,7 @@ class LogitsMetadata:
|
|||||||
draft_extend_select_index = None
|
draft_extend_select_index = None
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
forward_mode=forward_batch.forward_mode,
|
forward_mode=forward_mode,
|
||||||
capture_hidden_mode=forward_batch.capture_hidden_mode,
|
capture_hidden_mode=forward_batch.capture_hidden_mode,
|
||||||
next_token_logits_buffer=forward_batch.next_token_logits_buffer,
|
next_token_logits_buffer=forward_batch.next_token_logits_buffer,
|
||||||
extend_return_logprob=extend_return_logprob,
|
extend_return_logprob=extend_return_logprob,
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from unittest.mock import MagicMock
|
|||||||
|
|
||||||
import torch
|
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.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
@@ -40,6 +41,31 @@ def _logits_output(num_rows: int) -> SimpleNamespace:
|
|||||||
|
|
||||||
|
|
||||||
class TestMlpSyncPadUnpad(CustomTestCase):
|
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):
|
def test_init_mlp_sync_metadata_scales_speculative_request_width(self):
|
||||||
spec_info = SimpleNamespace(
|
spec_info = SimpleNamespace(
|
||||||
num_tokens_per_req=4,
|
num_tokens_per_req=4,
|
||||||
|
|||||||
Reference in New Issue
Block a user