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 functools import lru_cache
|
||||||
from http import HTTPStatus
|
from http import HTTPStatus
|
||||||
from itertools import chain
|
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 numpy as np
|
||||||
import torch
|
import torch
|
||||||
@@ -1379,6 +1389,12 @@ class Req(ReqDllmMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _MambaRadixCacheV2TrackEntry(NamedTuple):
|
||||||
|
track_mask: bool
|
||||||
|
track_index: int
|
||||||
|
track_seqlen: int
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||||
"""Store all information of a batch on the scheduler."""
|
"""Store all information of a batch on the scheduler."""
|
||||||
@@ -1850,12 +1866,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
req.is_retracted = False
|
req.is_retracted = False
|
||||||
|
|
||||||
if get_global_server_args().enable_mamba_extra_buffer():
|
if get_global_server_args().enable_mamba_extra_buffer():
|
||||||
self._mamba_radix_cache_v2_req_prepare_for_extend(
|
track_entry = self._mamba_radix_cache_v2_req_prepare_for_extend(req)
|
||||||
req,
|
mamba_track_mask_cpu.append(track_entry.track_mask)
|
||||||
mamba_track_mask_cpu,
|
mamba_track_indices_cpu.append(track_entry.track_index)
|
||||||
mamba_track_indices_cpu,
|
mamba_track_seqlens_cpu.append(track_entry.track_seqlen)
|
||||||
mamba_track_seqlens_cpu,
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.return_logprob:
|
if self.return_logprob:
|
||||||
# Find input logprob token ids.
|
# Find input logprob token ids.
|
||||||
@@ -2001,10 +2015,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
def _mamba_radix_cache_v2_req_prepare_for_extend(
|
def _mamba_radix_cache_v2_req_prepare_for_extend(
|
||||||
self,
|
self,
|
||||||
req: Req,
|
req: Req,
|
||||||
mamba_track_mask_cpu: List[bool],
|
) -> "_MambaRadixCacheV2TrackEntry":
|
||||||
mamba_track_indices_cpu: List[int],
|
|
||||||
mamba_track_seqlens_cpu: List[int],
|
|
||||||
):
|
|
||||||
def _force_track_h(i: int) -> int:
|
def _force_track_h(i: int) -> int:
|
||||||
assert i % FLA_CHUNK_SIZE == 0
|
assert i % FLA_CHUNK_SIZE == 0
|
||||||
# There are 3 cases for mamba_track_seqlen passed to mamba_track_seqlens_cpu:
|
# 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
|
mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size
|
||||||
mask = req.extend_input_len >= mamba_cache_chunk_size
|
mask = req.extend_input_len >= mamba_cache_chunk_size
|
||||||
mamba_track_mask_cpu.append(mask)
|
track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item()
|
||||||
mamba_track_indices_cpu.append(
|
|
||||||
req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item()
|
|
||||||
)
|
|
||||||
mamba_track_seqlen = -1
|
mamba_track_seqlen = -1
|
||||||
if mask:
|
if mask:
|
||||||
# mamba_track_seqlen is used to calculate the indices to track in
|
# 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 = _force_track_h(req.mamba_branching_seqlen)
|
||||||
mamba_track_seqlen_aligned = req.mamba_branching_seqlen
|
mamba_track_seqlen_aligned = req.mamba_branching_seqlen
|
||||||
req.mamba_last_track_seqlen = mamba_track_seqlen_aligned
|
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"):
|
def mix_with_running(self, running_batch: "ScheduleBatch"):
|
||||||
self.forward_mode = ForwardMode.MIXED
|
self.forward_mode = ForwardMode.MIXED
|
||||||
|
|||||||
Reference in New Issue
Block a user