Fix ScheduleBatch req pool CPU metadata (#28514)
This commit is contained in:
@@ -2596,6 +2596,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
def prepare_for_decode(self):
|
def prepare_for_decode(self):
|
||||||
self.forward_mode = ForwardMode.DECODE
|
self.forward_mode = ForwardMode.DECODE
|
||||||
bs = len(self.reqs)
|
bs = len(self.reqs)
|
||||||
|
if self.req_pool_indices_cpu is None and self.req_pool_indices is not None:
|
||||||
|
self.req_pool_indices_cpu = (
|
||||||
|
self.req_pool_indices.detach().cpu().to(dtype=torch.int64)
|
||||||
|
)
|
||||||
# Decode embeds the last output token via embed_tokens; clear the stale
|
# Decode embeds the last output token via embed_tokens; clear the stale
|
||||||
# prefill-time tensor so it doesn't leak into ForwardBatch.
|
# prefill-time tensor so it doesn't leak into ForwardBatch.
|
||||||
self.input_embeds = None
|
self.input_embeds = None
|
||||||
@@ -2690,6 +2694,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
if keep_indices is None or len(keep_indices) == 0:
|
if keep_indices is None or len(keep_indices) == 0:
|
||||||
# Filter out all requests
|
# Filter out all requests
|
||||||
self.reqs = []
|
self.reqs = []
|
||||||
|
self.req_pool_indices = torch.empty(
|
||||||
|
0, dtype=torch.int64, device=self.device
|
||||||
|
)
|
||||||
|
self.req_pool_indices_cpu = torch.empty(0, dtype=torch.int64)
|
||||||
|
self.seq_lens = torch.empty(0, dtype=torch.int64, device=self.device)
|
||||||
|
self.seq_lens_cpu = torch.empty(0, dtype=torch.int64)
|
||||||
|
self.orig_seq_lens = torch.empty(0, dtype=torch.int32, device=self.device)
|
||||||
|
self.out_cache_loc = None
|
||||||
|
self.seq_lens_sum = 0
|
||||||
return
|
return
|
||||||
|
|
||||||
if len(keep_indices) == len(self.reqs):
|
if len(keep_indices) == len(self.reqs):
|
||||||
@@ -2747,6 +2760,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def merge_batch(self, other: ScheduleBatch):
|
def merge_batch(self, other: ScheduleBatch):
|
||||||
|
if self.req_pool_indices_cpu is None and self.req_pool_indices is not None:
|
||||||
|
self.req_pool_indices_cpu = (
|
||||||
|
self.req_pool_indices.detach().cpu().to(dtype=torch.int64)
|
||||||
|
)
|
||||||
|
if other.req_pool_indices_cpu is None and other.req_pool_indices is not None:
|
||||||
|
other.req_pool_indices_cpu = (
|
||||||
|
other.req_pool_indices.detach().cpu().to(dtype=torch.int64)
|
||||||
|
)
|
||||||
|
|
||||||
# Penalizer orchestrator must be merged before Batch.reqs is merged. This is because
|
# Penalizer orchestrator must be merged before Batch.reqs is merged. This is because
|
||||||
# orchestrator.merge() depends on Batch.reqs during preparation of each penalizers, so it
|
# orchestrator.merge() depends on Batch.reqs during preparation of each penalizers, so it
|
||||||
# needs to be called with pre-merged Batch.reqs.
|
# needs to be called with pre-merged Batch.reqs.
|
||||||
|
|||||||
@@ -2438,9 +2438,11 @@ class Scheduler(
|
|||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
req_pool_indices = [r.req_pool_idx for r in reqs]
|
||||||
batch.req_pool_indices = torch.tensor(
|
batch.req_pool_indices = torch.tensor(
|
||||||
[r.req_pool_idx for r in reqs], dtype=torch.int64, device=device
|
req_pool_indices, dtype=torch.int64, device=device
|
||||||
)
|
)
|
||||||
|
batch.req_pool_indices_cpu = torch.tensor(req_pool_indices, dtype=torch.int64)
|
||||||
seq_lens = [len(r.origin_input_ids) + len(r.output_ids) - 1 for r in reqs]
|
seq_lens = [len(r.origin_input_ids) + len(r.output_ids) - 1 for r in reqs]
|
||||||
batch.seq_lens = torch.tensor(seq_lens, dtype=torch.int64, device=device)
|
batch.seq_lens = torch.tensor(seq_lens, dtype=torch.int64, device=device)
|
||||||
batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int64)
|
batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int64)
|
||||||
|
|||||||
@@ -0,0 +1,121 @@
|
|||||||
|
import types
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
|
from sglang.srt.managers.schedule_batch import ScheduleBatch # noqa: E402
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestScheduleBatchReqPoolIndices(unittest.TestCase):
|
||||||
|
def test_prepare_for_decode_restores_missing_req_pool_indices_cpu(self):
|
||||||
|
req = types.SimpleNamespace(
|
||||||
|
decode_batch_idx=0,
|
||||||
|
kv_committed_len=10,
|
||||||
|
kv_allocated_len=10,
|
||||||
|
)
|
||||||
|
batch = ScheduleBatch(
|
||||||
|
reqs=[req],
|
||||||
|
model_config=types.SimpleNamespace(is_encoder_decoder=False),
|
||||||
|
req_pool_indices=torch.tensor([4], dtype=torch.int64),
|
||||||
|
req_pool_indices_cpu=None,
|
||||||
|
seq_lens=torch.tensor([10], dtype=torch.int64),
|
||||||
|
seq_lens_cpu=torch.tensor([10], dtype=torch.int64),
|
||||||
|
orig_seq_lens=torch.tensor([10], dtype=torch.int32),
|
||||||
|
seq_lens_sum=10,
|
||||||
|
sampling_info=types.SimpleNamespace(
|
||||||
|
penalizer_orchestrator=types.SimpleNamespace(is_required=False)
|
||||||
|
),
|
||||||
|
spec_algorithm=types.SimpleNamespace(is_none=lambda: True),
|
||||||
|
enable_overlap=False,
|
||||||
|
device="cpu",
|
||||||
|
hisparse_coordinator=MagicMock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"sglang.srt.managers.schedule_batch.alloc_for_decode",
|
||||||
|
return_value=torch.tensor([42], dtype=torch.int64),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.managers.schedule_batch.get_global_server_args",
|
||||||
|
return_value=types.SimpleNamespace(
|
||||||
|
enable_mamba_extra_buffer=lambda: False
|
||||||
|
),
|
||||||
|
),
|
||||||
|
):
|
||||||
|
batch.prepare_for_decode()
|
||||||
|
|
||||||
|
self.assertTrue(torch.equal(batch.req_pool_indices_cpu, torch.tensor([4])))
|
||||||
|
batch.hisparse_coordinator.map_last_loc_to_buffer.assert_called_once()
|
||||||
|
|
||||||
|
def test_filter_batch_to_empty_clears_req_pool_metadata(self):
|
||||||
|
req = types.SimpleNamespace(finished=lambda: True)
|
||||||
|
batch = ScheduleBatch(
|
||||||
|
reqs=[req],
|
||||||
|
model_config=types.SimpleNamespace(is_encoder_decoder=False),
|
||||||
|
req_pool_indices=torch.tensor([4], dtype=torch.int64),
|
||||||
|
req_pool_indices_cpu=torch.tensor([4], dtype=torch.int64),
|
||||||
|
seq_lens=torch.tensor([10], dtype=torch.int64),
|
||||||
|
seq_lens_cpu=torch.tensor([10], dtype=torch.int64),
|
||||||
|
orig_seq_lens=torch.tensor([10], dtype=torch.int32),
|
||||||
|
seq_lens_sum=10,
|
||||||
|
device="cpu",
|
||||||
|
)
|
||||||
|
|
||||||
|
batch.filter_batch()
|
||||||
|
|
||||||
|
self.assertEqual(batch.req_pool_indices.numel(), 0)
|
||||||
|
self.assertEqual(batch.req_pool_indices_cpu.numel(), 0)
|
||||||
|
self.assertEqual(batch.seq_lens.numel(), 0)
|
||||||
|
self.assertEqual(batch.seq_lens_cpu.numel(), 0)
|
||||||
|
self.assertEqual(batch.seq_lens_sum, 0)
|
||||||
|
|
||||||
|
def test_merge_batch_restores_missing_req_pool_indices_cpu(self):
|
||||||
|
self_batch = ScheduleBatch(
|
||||||
|
reqs=[object(), object()],
|
||||||
|
model_config=types.SimpleNamespace(is_encoder_decoder=False),
|
||||||
|
req_pool_indices=torch.tensor([1, 2], dtype=torch.int64),
|
||||||
|
req_pool_indices_cpu=None,
|
||||||
|
seq_lens=torch.tensor([10, 20], dtype=torch.int64),
|
||||||
|
seq_lens_cpu=torch.tensor([10, 20], dtype=torch.int64),
|
||||||
|
orig_seq_lens=torch.tensor([10, 20], dtype=torch.int32),
|
||||||
|
seq_lens_sum=30,
|
||||||
|
sampling_info=MagicMock(),
|
||||||
|
return_logprob=False,
|
||||||
|
has_grammar=False,
|
||||||
|
return_hidden_states=False,
|
||||||
|
is_prefill_only=False,
|
||||||
|
)
|
||||||
|
other_batch = ScheduleBatch(
|
||||||
|
reqs=[object()],
|
||||||
|
model_config=types.SimpleNamespace(is_encoder_decoder=False),
|
||||||
|
req_pool_indices=torch.tensor([3], dtype=torch.int64),
|
||||||
|
req_pool_indices_cpu=None,
|
||||||
|
seq_lens=torch.tensor([30], dtype=torch.int64),
|
||||||
|
seq_lens_cpu=torch.tensor([30], dtype=torch.int64),
|
||||||
|
orig_seq_lens=torch.tensor([30], dtype=torch.int32),
|
||||||
|
seq_lens_sum=30,
|
||||||
|
sampling_info=MagicMock(),
|
||||||
|
return_logprob=False,
|
||||||
|
has_grammar=False,
|
||||||
|
return_hidden_states=False,
|
||||||
|
is_prefill_only=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self_batch.merge_batch(other_batch)
|
||||||
|
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(self_batch.req_pool_indices_cpu, torch.tensor([1, 2, 3]))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user