Return a mamba tracking entry from the cache lookup instead of mutating caller lists (#25724)

This commit is contained in:
fzyzcjy
2026-05-19 09:21:55 +08:00
committed by GitHub
parent 3fd6a58e6c
commit 2d2dff28da
+28 -16
View File
@@ -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