diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 4290a61e4..188956a8b 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -40,7 +40,17 @@ from enum import Enum, auto from functools import lru_cache from http import HTTPStatus from itertools import chain -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple, Union +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + NamedTuple, + Optional, + Set, + Tuple, + Union, +) import numpy as np import torch @@ -1379,6 +1389,12 @@ class Req(ReqDllmMixin): ) +class _MambaRadixCacheV2TrackEntry(NamedTuple): + track_mask: bool + track_index: int + track_seqlen: int + + @dataclasses.dataclass class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): """Store all information of a batch on the scheduler.""" @@ -1850,12 +1866,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): req.is_retracted = False if get_global_server_args().enable_mamba_extra_buffer(): - self._mamba_radix_cache_v2_req_prepare_for_extend( - req, - mamba_track_mask_cpu, - mamba_track_indices_cpu, - mamba_track_seqlens_cpu, - ) + track_entry = self._mamba_radix_cache_v2_req_prepare_for_extend(req) + mamba_track_mask_cpu.append(track_entry.track_mask) + mamba_track_indices_cpu.append(track_entry.track_index) + mamba_track_seqlens_cpu.append(track_entry.track_seqlen) if self.return_logprob: # Find input logprob token ids. @@ -2001,10 +2015,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): def _mamba_radix_cache_v2_req_prepare_for_extend( self, req: Req, - mamba_track_mask_cpu: List[bool], - mamba_track_indices_cpu: List[int], - mamba_track_seqlens_cpu: List[int], - ): + ) -> "_MambaRadixCacheV2TrackEntry": def _force_track_h(i: int) -> int: assert i % FLA_CHUNK_SIZE == 0 # There are 3 cases for mamba_track_seqlen passed to mamba_track_seqlens_cpu: @@ -2018,10 +2029,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size mask = req.extend_input_len >= mamba_cache_chunk_size - mamba_track_mask_cpu.append(mask) - mamba_track_indices_cpu.append( - req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item() - ) + track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item() mamba_track_seqlen = -1 if mask: # mamba_track_seqlen is used to calculate the indices to track in @@ -2076,7 +2084,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): mamba_track_seqlen = _force_track_h(req.mamba_branching_seqlen) mamba_track_seqlen_aligned = req.mamba_branching_seqlen req.mamba_last_track_seqlen = mamba_track_seqlen_aligned - mamba_track_seqlens_cpu.append(mamba_track_seqlen) + return _MambaRadixCacheV2TrackEntry( + track_mask=mask, + track_index=track_index, + track_seqlen=mamba_track_seqlen, + ) def mix_with_running(self, running_batch: "ScheduleBatch"): self.forward_mode = ForwardMode.MIXED