diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 02ff636cd..0ae504235 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2073,6 +2073,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): mamba_track_indices: torch.Tensor = None # shape: [b], int64 mamba_track_mask: torch.Tensor = None # shape: [b], bool mamba_track_seqlens: torch.Tensor = None # shape: [b], int64 + mamba_track_mask_cpu: Optional[List[bool]] = None # shape: [b] + mamba_track_mask_next_cpu: Optional[List[bool]] = None # shape: [b] + mamba_decode_batch_idx_cpu: Optional[List[int]] = None # shape: [b] # Lazy + spec: this iteration's per-req scatter positions # (see mamba_lazy_spec_prepare). mamba_lazy_spec_track_positions_cpu: Optional[List[int]] = None # shape: [b] @@ -3013,6 +3016,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # Spec decoding owns decode preparation (allocation, seq-lens bookkeeping). from sglang.srt.speculative.spec_utils import spec_prepare_for_decode + self.mamba_track_mask_cpu = None + self.mamba_track_mask_next_cpu = None + self.mamba_decode_batch_idx_cpu = None spec_prepare_for_decode(self) return @@ -3063,11 +3069,23 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.mamba_lazy_prealloc_at_boundary(mamba_track_interval) set_mamba_track_indices_from_reqs(self) + track_remainders_cpu = self.seq_lens_cpu % mamba_track_interval + track_mask_cpu = track_remainders_cpu == 0 + self.mamba_track_mask_cpu = track_mask_cpu.tolist() + self.mamba_track_mask_next_cpu = ( + (track_remainders_cpu == mamba_track_interval - 1).tolist() + if self.enable_overlap + else None + ) + # ScheduleBatch.copy() snapshots the list of requests, but the Req + # objects remain shared. The next overlapped decode can therefore + # advance their counters before this batch's result is processed. + self.mamba_decode_batch_idx_cpu = [ + req.decode_batch_idx for req in self.reqs + ] # async H2D - self.mamba_track_mask = ( - (self.seq_lens_cpu % mamba_track_interval == 0) - .pin_memory() - .to(device=self.device, non_blocking=True) + self.mamba_track_mask = track_mask_cpu.pin_memory().to( + device=self.device, non_blocking=True ) def filter_batch( @@ -3129,6 +3147,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.mamba_track_indices = None self.mamba_track_mask = None self.mamba_track_seqlens = None + self.mamba_track_mask_cpu = None + self.mamba_track_mask_next_cpu = None + self.mamba_decode_batch_idx_cpu = None self.mamba_lazy_spec_track_positions_cpu = None self.mamba_cow_src_indices = None self.mamba_cow_dst_indices = None @@ -3189,6 +3210,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.mamba_track_indices = None self.mamba_track_mask = None self.mamba_track_seqlens = None + self.mamba_track_mask_cpu = None + self.mamba_track_mask_next_cpu = None + self.mamba_decode_batch_idx_cpu = None self.mamba_lazy_spec_track_positions_cpu = None if self.return_logprob and other.return_logprob: self.top_logprobs_nums = self.top_logprobs_nums + other.top_logprobs_nums @@ -3247,6 +3271,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): mamba_track_indices=self.mamba_track_indices, mamba_track_mask=self.mamba_track_mask, mamba_track_seqlens=self.mamba_track_seqlens, + mamba_track_mask_cpu=self.mamba_track_mask_cpu, + mamba_track_mask_next_cpu=self.mamba_track_mask_next_cpu, + mamba_decode_batch_idx_cpu=self.mamba_decode_batch_idx_cpu, mamba_lazy_spec_track_positions_cpu=self.mamba_lazy_spec_track_positions_cpu, dp_cooperation_info=self.dp_cooperation_info, prefill_stats=self.prefill_stats, diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 24f0e089c..0f4ce352f 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -1015,9 +1015,28 @@ class SchedulerBatchResultProcessor: i: int, logits_output: LogitsProcessorOutput, ): + known_mamba_boundary = None + if batch.mamba_track_mask_cpu is not None: + lookahead = req.decode_batch_idx - batch.mamba_decode_batch_idx_cpu[i] + assert lookahead in (0, 1), ( + f"mamba result lookahead={lookahead} for req {req.rid}; " + "overlap advanced more than one decode batch" + ) + if lookahead == 0: + known_mamba_boundary = bool(batch.mamba_track_mask_cpu[i]) + else: + known_mamba_boundary = bool(batch.mamba_track_mask_next_cpu[i]) + # Called here (after update_finish_state) so req.finished() is valid # for mamba_lazy_post_decode_at_boundary inside. - self._mamba_prefix_cache_update(req, batch, result, i) + if known_mamba_boundary is None or known_mamba_boundary: + self._mamba_prefix_cache_update( + req, + batch, + result, + i, + known_boundary=known_mamba_boundary is True, + ) if ( get_disagg().disaggregation_decode_enable_offload_kvcache @@ -1078,6 +1097,7 @@ class SchedulerBatchResultProcessor: batch: ScheduleBatch, result: GenerationBatchResult, i: int, + known_boundary: bool = False, ) -> None: """Update mamba track state at ping-pong boundaries. @@ -1090,9 +1110,15 @@ class SchedulerBatchResultProcessor: return lazy = get_server_args().enable_mamba_extra_buffer_lazy() - at_boundary, track_seqlen = self._mamba_check_track_boundary( - req, batch, result, i - ) + if known_boundary: + self._mamba_assert_committed_len_lookahead(req) + track_seqlen = req.kv_committed_len + assert track_seqlen % get_exec().mamba.mamba_track_interval == 0 + at_boundary = True + else: + at_boundary, track_seqlen = self._mamba_check_track_boundary( + req, batch, result, i + ) if lazy and not batch.spec_algorithm.is_none(): # For spec, at_boundary means this step actually crossed an interval. @@ -1169,6 +1195,20 @@ class SchedulerBatchResultProcessor: # keep holds the track_seqlen state either way. req.mamba_last_track_seqlen = track_seqlen + @staticmethod + def _mamba_assert_committed_len_lookahead(req: Req) -> None: + """Alarm if overlap advances beyond the scheduler's modeled window.""" + assert req.output_ids, ( + "mamba track boundary reached before a decode token was appended " + f"(req {req.rid}); output_ids is empty" + ) + token_seq_len = len(req.origin_input_ids) + len(req.output_ids) - 1 + assert (req.kv_committed_len - token_seq_len) in (0, 1), ( + f"mamba track boundary: kv_committed_len={req.kv_committed_len} " + f"leads seq_len={token_seq_len} by more than one (req {req.rid}); " + "overlap lookahead wider than assumed" + ) + def _mamba_check_track_boundary(self, req, batch, result, i): """Check if this decode step crosses a mamba track interval boundary. @@ -1186,6 +1226,7 @@ class SchedulerBatchResultProcessor: interval = get_exec().mamba.mamba_track_interval if batch.spec_algorithm.is_none(): + self._mamba_assert_committed_len_lookahead(req) if req.kv_committed_len % interval == 0: return True, req.kv_committed_len elif result.num_correct_drafts_per_req_cpu is not None: diff --git a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py new file mode 100644 index 000000000..a40308602 --- /dev/null +++ b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py @@ -0,0 +1,199 @@ +import unittest +from array import array +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import torch + +from sglang.srt.managers.schedule_batch import Req, ScheduleBatch +from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.managers.scheduler_components.batch_result_processor import ( + SchedulerBatchResultProcessor, +) +from sglang.srt.sampling.sampling_params import SamplingParams +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +def _make_batch() -> tuple[Req, ScheduleBatch]: + sampling_params = SamplingParams(max_new_tokens=32) + sampling_params.normalize(None) + req = Req( + rid="req", + origin_input_text="", + origin_input_ids=array("q", [1, 2]), + sampling_params=sampling_params, + vocab_size=128, + ) + req.output_ids.append(3) + req.kv_committed_len = 2 + + batch = ScheduleBatch(reqs=[req]) + batch.device = "cpu" + batch.model_config = SimpleNamespace(is_encoder_decoder=False) + batch.enable_overlap = True + batch.spec_algorithm = SimpleNamespace(is_none=lambda: True) + batch.sampling_info = SimpleNamespace( + penalizer_orchestrator=SimpleNamespace(is_required=False) + ) + batch.hisparse_coordinator = None + batch.seq_lens = torch.tensor([2], dtype=torch.int64) + batch.seq_lens_cpu = torch.tensor([2], dtype=torch.int64) + batch.orig_seq_lens = torch.tensor([2], dtype=torch.int32) + return req, batch + + +def _make_processor() -> SchedulerBatchResultProcessor: + metrics_reporter = MagicMock() + metrics_reporter.num_generated_tokens = 0 + metrics_reporter.forward_ct_decode = 0 + return SchedulerBatchResultProcessor( + is_generation=True, + disaggregation_mode=None, + enable_overlap=True, + enable_overlap_mlx=False, + server_args=SimpleNamespace(), + model_config=SimpleNamespace(think_end_ids=None), + token_to_kv_pool_allocator=MagicMock(), + tree_cache=None, + hisparse_coordinator=None, + req_to_token_pool=None, + decode_offload_manager=None, + metrics_collector=None, + metrics_reporter=metrics_reporter, + draft_worker=None, + model_worker=MagicMock(), + logprob_result_processor=None, + output_streamer=MagicMock(), + abort_request=lambda *args, **kwargs: None, + ) + + +def _make_result(): + return SimpleNamespace( + copy_done=None, + routed_experts_output=None, + indexer_topk_output=None, + logits_output=SimpleNamespace(hidden_states=None, customized_info=None), + next_token_ids=[4], + can_run_cuda_graph=False, + num_correct_drafts=0, + num_block_accept_tokens=0, + num_cap_tokens=0, + speculative_num_draft_tokens=0, + ) + + +class TestMambaBoundaryMaskReuse(unittest.TestCase): + def test_overlap_scheduler_handles_zero_and_one_batch_lookahead(self): + for schedule_next_decode, expected_lookahead in ((False, 0), (True, 1)): + with self.subTest(schedule_next_decode=schedule_next_decode): + req, batch = _make_batch() + processor = _make_processor() + result = _make_result() + + scheduler = Scheduler.__new__(Scheduler) + scheduler.gracefully_exit = False + scheduler.request_receiver = MagicMock() + scheduler.request_receiver.recv_requests.side_effect = [ + [], + [], + StopIteration, + ] + scheduler.process_input_requests = MagicMock() + scheduler._engine_paused = False + scheduler.running_batch = batch + scheduler.is_disable_overlap_for_batch = MagicMock(return_value=False) + scheduler.run_batch = MagicMock(return_value=result) + scheduler._apply_war_barrier = MagicMock() + scheduler.is_generation = False + scheduler.last_batch = None + + plan_count = 0 + + def get_next_batch_to_run(*, running_batch, last_batch): + nonlocal plan_count + del running_batch, last_batch + plan_count += 1 + if plan_count == 1: + batch.prepare_for_decode() + return SimpleNamespace( + running_batch=batch, + batch_to_run=batch, + ) + if plan_count == 2 and schedule_next_decode: + batch.prepare_for_decode() + return SimpleNamespace( + running_batch=batch, + batch_to_run=batch, + ) + return SimpleNamespace( + running_batch=batch, + batch_to_run=None, + ) + + scheduler.get_next_batch_to_run = get_next_batch_to_run + observed_lookahead = [] + + def process_batch_result(result_batch, batch_result): + observed_lookahead.append( + req.decode_batch_idx + - result_batch.mamba_decode_batch_idx_cpu[0] + ) + processor.process_batch_result_decode(result_batch, batch_result) + + scheduler.process_batch_result = process_batch_result + server_args = SimpleNamespace( + enable_mamba_extra_buffer=lambda: True, + enable_mamba_extra_buffer_lazy=lambda: False, + ) + + with ( + patch( + "sglang.srt.managers.schedule_batch.alloc_for_decode", + return_value=torch.tensor([3], dtype=torch.int64), + ), + patch( + "sglang.srt.managers.schedule_batch.get_server_args", + return_value=server_args, + ), + patch( + "sglang.srt.managers.schedule_batch.get_exec", + return_value=SimpleNamespace( + mamba=SimpleNamespace(mamba_track_interval=4) + ), + ), + patch( + "sglang.srt.managers.schedule_batch.set_mamba_track_indices_from_reqs" + ), + patch.object(torch.Tensor, "pin_memory", lambda tensor: tensor), + patch( + "sglang.srt.managers.scheduler_components." + "batch_result_processor.get_observability", + return_value=SimpleNamespace(enable_metrics=False), + ), + patch( + "sglang.srt.managers.scheduler_components." + "batch_result_processor.get_disagg", + return_value=SimpleNamespace( + disaggregation_decode_enable_offload_kvcache=False + ), + ), + patch.object( + SchedulerBatchResultProcessor, + "_mamba_prefix_cache_update", + ) as cache_update, + ): + with self.assertRaises(StopIteration): + scheduler.event_loop_overlap() + + self.assertEqual(observed_lookahead, [expected_lookahead]) + if expected_lookahead == 0: + cache_update.assert_not_called() + else: + self.assertTrue(cache_update.call_args.kwargs["known_boundary"]) + + +if __name__ == "__main__": + unittest.main()