[Sampling] Support sampling masks with overlap scheduling (#36631)
Co-authored-by: ByronHsu <ByronHsu@users.noreply.github.com> Co-authored-by: root <root@slurm-h200-208-179.slurm-compute.tenant-slurm.svc.cluster.local> Co-authored-by: Byron Hsu <byron+per@periodiclabs.ai>
This commit is contained in:
co-authored by
ByronHsu
root
Byron Hsu
parent
55b45cb45a
commit
fd7743e0e1
@@ -46,6 +46,7 @@ from sglang.srt.speculative.eagle_disaggregation import (
|
||||
build_eagle_disagg_draft_input,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
|
||||
|
||||
@@ -335,9 +336,14 @@ class TestMooncakePPStaging(unittest.TestCase):
|
||||
)
|
||||
|
||||
|
||||
class TestEagleDsaSeedTransfer(unittest.TestCase):
|
||||
class TestEagleDsaSeedTransfer(CustomTestCase):
|
||||
@staticmethod
|
||||
def _make_req(seed, metadata_buffer_index=0):
|
||||
def _make_req(
|
||||
seed,
|
||||
metadata_buffer_index=0,
|
||||
sampling_mask=None,
|
||||
sampling_logprob=None,
|
||||
):
|
||||
return SimpleNamespace(
|
||||
metadata_buffer_index=metadata_buffer_index,
|
||||
output_ids=[101],
|
||||
@@ -347,7 +353,13 @@ class TestEagleDsaSeedTransfer(unittest.TestCase):
|
||||
cached_tokens_storage=0,
|
||||
multimodal_inputs=None,
|
||||
return_logprob=False,
|
||||
return_sampling_mask=False,
|
||||
return_sampling_mask=sampling_mask is not None,
|
||||
output_token_sampling_mask=(
|
||||
None if sampling_mask is None else [sampling_mask]
|
||||
),
|
||||
output_token_sampling_logprobs=(
|
||||
None if sampling_logprob is None else [sampling_logprob]
|
||||
),
|
||||
hidden_states_tensor=torch.tensor([1.0, 2.0]),
|
||||
output_topk_p=torch.tensor([1.0]),
|
||||
output_topk_index=torch.tensor([7]),
|
||||
@@ -360,6 +372,7 @@ class TestEagleDsaSeedTransfer(unittest.TestCase):
|
||||
size=2,
|
||||
hidden_size=2,
|
||||
hidden_states_dtype=torch.float32,
|
||||
max_sampling_mask_tokens=16,
|
||||
output_dsa_topk_indices_dim=3,
|
||||
)
|
||||
seed = torch.tensor([4, 5, 6], dtype=torch.int32)
|
||||
@@ -378,6 +391,46 @@ class TestEagleDsaSeedTransfer(unittest.TestCase):
|
||||
self.assertEqual(data_lens[-2], buffers.output_dsa_topk_indices.nbytes)
|
||||
self.assertEqual(item_lens[-2], buffers.output_dsa_topk_indices[0].nbytes)
|
||||
|
||||
def test_sampling_mask_metadata_is_opt_in(self):
|
||||
"""Disabled masks stay off the wire; enabled masks round-trip at capacity."""
|
||||
schemas = []
|
||||
for enabled in (False, True):
|
||||
with (
|
||||
self.subTest(enabled=enabled),
|
||||
envs.SGLANG_ENABLE_DISAGG_SAMPLING_MASK.override(enabled),
|
||||
):
|
||||
buffers = MetadataBuffers(
|
||||
size=1,
|
||||
hidden_size=2,
|
||||
hidden_states_dtype=torch.float32,
|
||||
max_sampling_mask_tokens=3,
|
||||
)
|
||||
buffers.set_buf(
|
||||
self._make_req(
|
||||
None,
|
||||
sampling_mask=[7, 8, 9] if enabled else None,
|
||||
sampling_logprob=-1.25 if enabled else None,
|
||||
)
|
||||
)
|
||||
schemas.append(buffers.get_buf_infos())
|
||||
if enabled:
|
||||
self.assertEqual(
|
||||
buffers.output_token_sampling_mask_idx.shape, (1, 3)
|
||||
)
|
||||
length, mask, logprob = buffers.get_buf(0)[6:9]
|
||||
self.assertEqual(length[0].item(), 3)
|
||||
self.assertEqual(mask.tolist(), [7, 8, 9])
|
||||
self.assertAlmostEqual(logprob[0].item(), -1.25)
|
||||
else:
|
||||
self.assertIsNone(buffers.output_token_sampling_mask_len)
|
||||
self.assertIsNone(buffers.output_token_sampling_mask_idx)
|
||||
self.assertIsNone(buffers.output_token_sampling_logprobs)
|
||||
self.assertEqual(buffers.get_buf(0)[6:9], (None, None, None))
|
||||
disabled_ptrs, _, disabled_sizes = schemas[0]
|
||||
enabled_ptrs, _, enabled_sizes = schemas[1]
|
||||
self.assertEqual(len(enabled_ptrs) - len(disabled_ptrs), 3)
|
||||
self.assertEqual(sum(enabled_sizes) - sum(disabled_sizes), 3 * 4 + 128)
|
||||
|
||||
def test_decode_input_requires_valid_seed_for_every_request(self):
|
||||
seeds = (
|
||||
torch.tensor([1, 2, 3], dtype=torch.int32),
|
||||
|
||||
@@ -5,8 +5,13 @@ import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.disaggregation.prefill import SchedulerDisaggregationPrefillMixin
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput, SamplingMaskStatus
|
||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT, ReqKvInfo
|
||||
from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
||||
SchedulerBatchResultProcessor,
|
||||
)
|
||||
from sglang.srt.managers.utils import GenerationBatchResult
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
@@ -205,5 +210,54 @@ def test_aborted_result_releases_mamba_allocated_before_kv():
|
||||
scheduler.output_streamer.stream_output.assert_called_once_with([req], False)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("transport_error", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
"status,http_status,err_type",
|
||||
[
|
||||
(SamplingMaskStatus.OVERFLOW, 400, "BadRequestError"),
|
||||
(SamplingMaskStatus.INVALID, 500, "InternalServerError"),
|
||||
],
|
||||
)
|
||||
@patch("sglang.srt.disaggregation.prefill.release_kv_cache", side_effect=_free_req)
|
||||
def test_sampling_mask_abort_preserves_error_and_releases_once(
|
||||
release_kv_cache, status, http_status, err_type, transport_error
|
||||
):
|
||||
"""A failed sender notification must not leak ownership or lose the API error."""
|
||||
scheduler = _Scheduler()
|
||||
scheduler.batch_result_processor.get_sampling_mask_finish_reason = lambda **kwargs: (
|
||||
SchedulerBatchResultProcessor.get_sampling_mask_finish_reason(None, **kwargs)
|
||||
)
|
||||
req = _Req(inflight_middle_chunks=0)
|
||||
req.to_finish = None
|
||||
req.return_sampling_mask = True
|
||||
req.time_stats.trace_ctx = Mock()
|
||||
if transport_error:
|
||||
req.disagg_kv_sender.abort.side_effect = RuntimeError("transport is down")
|
||||
result = GenerationBatchResult(
|
||||
next_token_ids=torch.tensor([11]),
|
||||
logits_output=LogitsProcessorOutput(
|
||||
next_token_logits=None, next_token_sampling_mask_status=[status]
|
||||
),
|
||||
)
|
||||
|
||||
with get_context().override_server_args(sampling_mask_max_tokens=64):
|
||||
scheduler.process_batch_result_disagg_prefill(_batch(req), result)
|
||||
scheduler.process_batch_result_disagg_prefill(_batch(req), result)
|
||||
|
||||
assert req.finished_reason.status_code == http_status
|
||||
assert req.finished_reason.err_type == err_type
|
||||
assert req.output_ids == []
|
||||
assert not req.kv.holds_kv and not req.kv.holds_mamba
|
||||
assert req.metadata_buffer_index == -1
|
||||
assert not req.pending_bootstrap
|
||||
assert req.rid not in scheduler.disagg_prefill_pending_chunk_rids
|
||||
release_kv_cache.assert_called_once_with(req, scheduler.tree_cache, is_insert=False)
|
||||
req.disagg_kv_sender.abort.assert_called_once_with()
|
||||
scheduler.req_to_metadata_buffer_idx_allocator.free.assert_called_once_with(7)
|
||||
scheduler.tree_cache.release_aborted_request.assert_called_once_with(req.rid)
|
||||
scheduler.output_streamer.stream_output.assert_called_once_with([req], False)
|
||||
scheduler.send_kv_chunk.assert_not_called()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(pytest.main([__file__, "-v"]))
|
||||
|
||||
@@ -9,6 +9,7 @@ Requires: torch, sglang (run in an environment with sglang installed)
|
||||
|
||||
import gc
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
from weakref import WeakKeyDictionary as WeakKeyDict
|
||||
|
||||
@@ -20,9 +21,14 @@ from sglang.srt.disaggregation.decode_kvcache_offload_manager import (
|
||||
from sglang.srt.disaggregation.kv_events import OffloadedState
|
||||
from sglang.srt.managers.cache_controller import HiCacheAck
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
||||
SchedulerBatchResultProcessor,
|
||||
)
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
|
||||
@@ -440,5 +446,46 @@ class TestReleaseFinishedReq(unittest.TestCase):
|
||||
self.assertEqual(len(manager.offload_inflight), 0)
|
||||
|
||||
|
||||
class TestSamplingMaskAbortOffload(CustomTestCase):
|
||||
def test_abort_waits_for_existing_offload_before_reusing_slots(self):
|
||||
"""An abort must not recycle slots while a previous D2H copy reads them."""
|
||||
for inflight in (False, True):
|
||||
with self.subTest(inflight=inflight):
|
||||
manager, freed = _make_manager(pool_size=32)
|
||||
req = _make_mock_req(0, 20, 20)
|
||||
req.multimodal_inputs = None
|
||||
req.finished.return_value = True
|
||||
manager.req_to_token_pool.free.side_effect = lambda req: setattr(
|
||||
req.kv, "req_pool_idx", None
|
||||
)
|
||||
processor = SimpleNamespace(decode_offload_manager=manager)
|
||||
if inflight:
|
||||
manager.offload_inflight[req] = 1
|
||||
manager.ongoing_offload[1] = (req, torch.arange(4), [1], 0.0)
|
||||
manager.cache_controller = MagicMock()
|
||||
manager.cache_controller.ack_write_queue = [
|
||||
HiCacheAck(None, _FinishedEvent(), [1])
|
||||
]
|
||||
manager._trigger_backup = MagicMock(return_value="hash")
|
||||
|
||||
with get_context().override_server_args(
|
||||
disaggregation_decode_enable_offload_kvcache=True,
|
||||
enable_hisparse=False,
|
||||
):
|
||||
SchedulerBatchResultProcessor._handle_sampling_mask_abort(
|
||||
processor, req
|
||||
)
|
||||
|
||||
if inflight:
|
||||
self.assertEqual(freed, [])
|
||||
self.assertEqual(req.kv.req_pool_idx, 0)
|
||||
manager._check_offload_progress(1)
|
||||
self.assertEqual(len(freed), 1)
|
||||
self.assertTrue(torch.equal(freed[0], torch.arange(20)))
|
||||
self.assertIsNone(req.kv.req_pool_idx)
|
||||
manager.finalize_release_on_finish(req)
|
||||
self.assertEqual(len(freed), 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user