Return a mamba tracking entry from the cache lookup instead of mutating caller lists (#25724)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user