[Qwen3-next] support mamba radix cache for overlap scheduler (#14792)
This commit is contained in:
@@ -57,6 +57,7 @@ from sglang.srt.disaggregation.decode_schedule_batch_mixin import (
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
|
||||
from sglang.srt.mem_cache.allocator import (
|
||||
BaseTokenToKVPoolAllocator,
|
||||
SWATokenToKVPoolAllocator,
|
||||
@@ -543,6 +544,14 @@ class Req:
|
||||
# Memory pool info
|
||||
self.req_pool_idx: Optional[int] = None
|
||||
self.mamba_pool_idx: Optional[torch.Tensor] = None # shape (1)
|
||||
self.mamba_ping_pong_track_buffer: Optional[torch.Tensor] = None # shape (2)
|
||||
self.mamba_next_track_idx: Optional[int] = None # 0 or 1
|
||||
self.mamba_last_track_seqlen: Optional[int] = (
|
||||
None # seq len of the last cached mamba state
|
||||
)
|
||||
# the branching point seqlen to track mamba state. If set, given by prefix match,
|
||||
# it will be the tracked seqlen in the ping pong buffer for the right prefill pass.
|
||||
self.mamba_branching_seqlen: Optional[int] = None
|
||||
|
||||
# Check finish
|
||||
self.tokenizer = None
|
||||
@@ -824,11 +833,13 @@ class Req:
|
||||
self.last_node,
|
||||
self.last_host_node,
|
||||
self.host_hit_length,
|
||||
self.mamba_branching_seqlen,
|
||||
) = (
|
||||
match_result.device_indices,
|
||||
match_result.last_device_node,
|
||||
match_result.last_host_node,
|
||||
match_result.host_hit_length,
|
||||
match_result.mamba_branching_seqlen,
|
||||
)
|
||||
self.cache_protected_len = len(self.prefix_indices)
|
||||
|
||||
@@ -1027,6 +1038,10 @@ class Req:
|
||||
self.extend_logprob_start_len = 0
|
||||
self.is_chunked = 0
|
||||
self.mamba_pool_idx = None
|
||||
self.mamba_ping_pong_track_buffer = None
|
||||
self.mamba_next_track_idx = None
|
||||
self.mamba_last_track_seqlen = None
|
||||
self.mamba_branching_seqlen = None
|
||||
self.already_computed = 0
|
||||
self.kv_allocated_len = 0
|
||||
self.kv_committed_len = 0
|
||||
@@ -1115,6 +1130,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
out_cache_loc: torch.Tensor = None # shape: [b], int64
|
||||
output_ids: torch.Tensor = None # shape: [b], int64
|
||||
|
||||
# For hybrid GDN prefix cache
|
||||
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
|
||||
|
||||
# For multimodal inputs
|
||||
multimodal_inputs: Optional[List] = None
|
||||
|
||||
@@ -1380,6 +1400,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
input_embeds = []
|
||||
extend_input_logprob_token_ids = []
|
||||
multimodal_inputs = []
|
||||
mamba_track_mask_cpu = []
|
||||
mamba_track_indices_cpu = []
|
||||
mamba_track_seqlens_cpu = []
|
||||
|
||||
for i, (req, seq_len, pre_len) in enumerate(zip(reqs, seq_lens, prefix_lens)):
|
||||
req.req_pool_idx = req_pool_indices[i]
|
||||
@@ -1403,6 +1426,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
req.already_computed = seq_len
|
||||
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,
|
||||
)
|
||||
|
||||
# Compute the relative logprob_start_len in an extend batch
|
||||
#
|
||||
# Key variables:
|
||||
@@ -1512,6 +1543,23 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.extend_logprob_start_lens = [r.extend_logprob_start_len for r in reqs]
|
||||
self.extend_input_logprob_token_ids = extend_input_logprob_token_ids
|
||||
|
||||
if get_global_server_args().enable_mamba_extra_buffer():
|
||||
self.mamba_track_indices = torch.tensor(
|
||||
mamba_track_indices_cpu,
|
||||
dtype=torch.int64,
|
||||
device=self.device,
|
||||
)
|
||||
self.mamba_track_mask = torch.tensor(
|
||||
mamba_track_mask_cpu,
|
||||
dtype=torch.bool,
|
||||
device=self.device,
|
||||
)
|
||||
self.mamba_track_seqlens = torch.tensor(
|
||||
mamba_track_seqlens_cpu,
|
||||
dtype=torch.int64,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
if self.model_config.is_encoder_decoder:
|
||||
self.prepare_encoder_info_extend(input_ids, seq_lens)
|
||||
|
||||
@@ -1521,6 +1569,60 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.model_config.vocab_size,
|
||||
)
|
||||
|
||||
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],
|
||||
):
|
||||
mask = (req.extend_input_len // FLA_CHUNK_SIZE) * FLA_CHUNK_SIZE > 0
|
||||
mamba_track_mask_cpu.append(mask)
|
||||
mamba_track_indices_cpu.append(
|
||||
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
|
||||
# hybrid_linear_attn_backend's _init_track_ssm_indices. Due to the
|
||||
# fact that the ssm state between aligned and non-aligned are retrieved differently,
|
||||
# if 1) last pos and 2) is aligned, then retrieved from the last_recurrent_state,
|
||||
# otherwise retrieved from h (i.e. unaligned).
|
||||
# We need to pass the non-aligned seqlen to the calculation. Even though
|
||||
# we pass in mamba_track_seqlen, the actual tracked seqlen is mamba_last_track_seqlen.
|
||||
mamba_track_seqlen = len(req.prefix_indices) + req.extend_input_len
|
||||
# mamba_last_track_seqlen is actual tracked seqlen. Used to pass to
|
||||
# mamba radix cache to track which seqlen this mamba state should store at.
|
||||
mamba_track_seqlen_aligned = (
|
||||
len(req.prefix_indices)
|
||||
+ (req.extend_input_len // FLA_CHUNK_SIZE) * FLA_CHUNK_SIZE
|
||||
)
|
||||
req.mamba_next_track_idx = (
|
||||
self.req_to_token_pool.get_mamba_ping_pong_other_idx(
|
||||
req.mamba_next_track_idx
|
||||
)
|
||||
)
|
||||
if req.mamba_branching_seqlen is not None:
|
||||
# track branching point in this forward if the branching point
|
||||
# is within the current extend batch.
|
||||
branching_seqlen_aligned_mask = (
|
||||
req.mamba_branching_seqlen - len(req.prefix_indices)
|
||||
) % FLA_CHUNK_SIZE == 0
|
||||
if (
|
||||
req.mamba_branching_seqlen > len(req.prefix_indices)
|
||||
and req.mamba_branching_seqlen < mamba_track_seqlen
|
||||
and branching_seqlen_aligned_mask
|
||||
):
|
||||
# NOTE: See the comment above for mamba_track_seqlen, the +1 is necessary
|
||||
# because the branching point is not the last aligned position, so we need
|
||||
# to retrieve its state from h. Adding 1 will give us the correct index in h,
|
||||
# otherwise the calculation will retrieve the state from the last_recurrent_state,
|
||||
# which is not correct.
|
||||
mamba_track_seqlen = req.mamba_branching_seqlen + 1
|
||||
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)
|
||||
|
||||
def prepare_for_split_prefill(self):
|
||||
self.prepare_for_extend()
|
||||
# For split prefill, we need to set the forward mode to SPLIT_PREFILL
|
||||
@@ -1786,6 +1888,24 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.orig_seq_lens.add_(1)
|
||||
self.seq_lens_sum += bs
|
||||
|
||||
if get_global_server_args().enable_mamba_extra_buffer():
|
||||
self.mamba_track_indices = torch.tensor(
|
||||
[
|
||||
req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx]
|
||||
for req in self.reqs
|
||||
],
|
||||
dtype=torch.int64,
|
||||
device=self.device,
|
||||
)
|
||||
self.mamba_track_mask = torch.tensor(
|
||||
[
|
||||
sl % get_global_server_args().mamba_track_interval == 0
|
||||
for sl in self.seq_lens_cpu
|
||||
],
|
||||
dtype=torch.bool,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def maybe_wait_verify_done(self):
|
||||
if self.is_v2_eagle:
|
||||
draft_input: EagleDraftInput = self.spec_info
|
||||
@@ -1842,6 +1962,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.out_cache_loc = None
|
||||
self.seq_lens_sum = self.seq_lens.sum().item()
|
||||
self.output_ids = self.output_ids[keep_indices_device]
|
||||
self.mamba_track_indices = None
|
||||
self.mamba_track_mask = None
|
||||
self.mamba_track_seqlens = None
|
||||
self.return_logprob = any(req.return_logprob for req in self.reqs)
|
||||
if self.return_logprob:
|
||||
self.top_logprobs_nums = [self.top_logprobs_nums[i] for i in keep_indices]
|
||||
@@ -1889,6 +2012,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
self.seq_lens_sum += other.seq_lens_sum
|
||||
if self.output_ids is not None:
|
||||
self.output_ids = torch.cat([self.output_ids, other.output_ids])
|
||||
self.mamba_track_indices = None
|
||||
self.mamba_track_mask = None
|
||||
self.mamba_track_seqlens = None
|
||||
if self.return_logprob and other.return_logprob:
|
||||
self.top_logprobs_nums.extend(other.top_logprobs_nums)
|
||||
self.token_ids_logprobs.extend(other.token_ids_logprobs)
|
||||
@@ -1982,6 +2108,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
dllm_config=self.dllm_config,
|
||||
reqs=self.reqs,
|
||||
has_grammar=self.has_grammar,
|
||||
mamba_track_indices=self.mamba_track_indices,
|
||||
mamba_track_mask=self.mamba_track_mask,
|
||||
mamba_track_seqlens=self.mamba_track_seqlens,
|
||||
)
|
||||
|
||||
def copy(self):
|
||||
@@ -2003,6 +2132,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
is_prefill_only=self.is_prefill_only,
|
||||
seq_lens_cpu=self.seq_lens_cpu,
|
||||
enable_overlap=self.enable_overlap,
|
||||
mamba_track_indices=self.mamba_track_indices,
|
||||
mamba_track_mask=self.mamba_track_mask,
|
||||
mamba_track_seqlens=self.mamba_track_seqlens,
|
||||
)
|
||||
|
||||
def _is_available_size_sufficient(self, num_tokens: int) -> bool:
|
||||
@@ -2104,3 +2236,8 @@ class ModelWorkerBatch:
|
||||
# FIXME(lsyin): remove this after fully overlap grammar
|
||||
reqs: Optional[List[Req]] = None
|
||||
has_grammar: bool = False
|
||||
|
||||
# For mamba state tracking
|
||||
mamba_track_indices: Optional[torch.Tensor] = None # shape: [b], int64
|
||||
mamba_track_mask: Optional[torch.Tensor] = None # shape: [b], bool
|
||||
mamba_track_seqlens: Optional[torch.Tensor] = None # shape: [b], int64
|
||||
|
||||
@@ -363,7 +363,7 @@ class PrefillAdder:
|
||||
self.is_hybrid_swa = isinstance(
|
||||
self.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator
|
||||
)
|
||||
self.is_ssm_radix_cache = isinstance(self.tree_cache, MambaRadixCache)
|
||||
self.is_hybrid_ssm_cache = isinstance(self.tree_cache, MambaRadixCache)
|
||||
|
||||
self.priority_scheduling_preemption_threshold = (
|
||||
priority_scheduling_preemption_threshold
|
||||
@@ -389,7 +389,7 @@ class PrefillAdder:
|
||||
self.token_to_kv_pool_allocator.swa_available_size()
|
||||
+ self.tree_cache.swa_evictable_size(),
|
||||
)
|
||||
elif self.is_ssm_radix_cache:
|
||||
elif self.is_hybrid_ssm_cache:
|
||||
available_and_evictable = (
|
||||
self.token_to_kv_pool_allocator.available_size()
|
||||
+ self.tree_cache.full_evictable_size()
|
||||
@@ -411,7 +411,7 @@ class PrefillAdder:
|
||||
self.token_to_kv_pool_allocator.swa_available_size()
|
||||
+ self.tree_cache.swa_evictable_size(),
|
||||
)
|
||||
elif self.is_ssm_radix_cache:
|
||||
elif self.is_hybrid_ssm_cache:
|
||||
available_and_evictable = (
|
||||
self.token_to_kv_pool_allocator.available_size()
|
||||
+ self.tree_cache.full_evictable_size()
|
||||
|
||||
@@ -405,7 +405,7 @@ class Scheduler(
|
||||
|
||||
# Hybrid memory pool
|
||||
self.is_hybrid_swa = self.tp_worker.is_hybrid_swa
|
||||
self.is_ssm_model = (
|
||||
self.is_hybrid_ssm = (
|
||||
self.tp_worker.model_runner.hybrid_gdn_config is not None
|
||||
or self.tp_worker.model_runner.mamba2_config is not None
|
||||
)
|
||||
@@ -772,6 +772,7 @@ class Scheduler(
|
||||
eviction_policy=server_args.radix_eviction_policy,
|
||||
enable_metrics=self.enable_metrics,
|
||||
enable_kv_cache_events=self.enable_kv_cache_events,
|
||||
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
|
||||
)
|
||||
|
||||
if (
|
||||
@@ -808,7 +809,7 @@ class Scheduler(
|
||||
self.tree_cache = SWARadixCache(
|
||||
params=params, sliding_window_size=self.sliding_window_size
|
||||
)
|
||||
elif self.is_ssm_model:
|
||||
elif self.is_hybrid_ssm:
|
||||
from sglang.srt.mem_cache.mamba_radix_cache import MambaRadixCache
|
||||
|
||||
self.tree_cache = MambaRadixCache(params)
|
||||
|
||||
@@ -112,7 +112,7 @@ class SchedulerMetricsMixin:
|
||||
f"full token usage: {full_token_usage:.2f}, "
|
||||
f"swa token usage: {swa_token_usage:.2f}, "
|
||||
)
|
||||
elif self.is_ssm_model:
|
||||
elif self.is_hybrid_ssm:
|
||||
(
|
||||
full_num_used,
|
||||
_,
|
||||
@@ -166,7 +166,7 @@ class SchedulerMetricsMixin:
|
||||
self.stats.token_usage = token_usage
|
||||
if self.is_hybrid_swa:
|
||||
self.stats.swa_token_usage = swa_token_usage
|
||||
if self.is_ssm_model:
|
||||
if self.is_hybrid_ssm:
|
||||
self.stats.mamba_usage = mamba_usage
|
||||
self.stats.num_queue_reqs = len(self.waiting_queue)
|
||||
self.stats.num_grammar_queue_reqs = len(self.grammar_queue)
|
||||
@@ -238,7 +238,7 @@ class SchedulerMetricsMixin:
|
||||
f"#swa token: {swa_num_used}, "
|
||||
f"swa token usage: {swa_token_usage:.2f}, "
|
||||
)
|
||||
elif self.is_ssm_model:
|
||||
elif self.is_hybrid_ssm:
|
||||
(
|
||||
full_num_used,
|
||||
mamba_used,
|
||||
@@ -315,7 +315,7 @@ class SchedulerMetricsMixin:
|
||||
self.stats.token_usage = token_usage
|
||||
if self.is_hybrid_swa:
|
||||
self.stats.swa_token_usage = swa_token_usage
|
||||
if self.is_ssm_model:
|
||||
if self.is_hybrid_ssm:
|
||||
self.stats.mamba_usage = mamba_usage
|
||||
self.stats.decode_sum_seq_lens = batch.seq_lens_cpu.sum().item()
|
||||
self.stats.gen_throughput = self.last_gen_throughput
|
||||
@@ -402,7 +402,7 @@ class SchedulerMetricsMixin:
|
||||
if self.is_hybrid_swa:
|
||||
full_num_used, swa_num_used, *_ = self._get_swa_token_info()
|
||||
num_tokens = max(full_num_used, swa_num_used)
|
||||
elif self.is_ssm_model:
|
||||
elif self.is_hybrid_ssm:
|
||||
num_tokens = self._get_mamba_token_info()[0]
|
||||
else:
|
||||
num_tokens = self._get_token_info()[0]
|
||||
|
||||
@@ -21,6 +21,7 @@ from sglang.srt.managers.schedule_batch import (
|
||||
ScheduleBatch,
|
||||
)
|
||||
from sglang.srt.mem_cache.common import release_kv_cache
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.tracing.trace import trace_slice, trace_slice_batch, trace_slice_end
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -267,6 +268,7 @@ class SchedulerOutputProcessorMixin:
|
||||
next_token_ids = result.next_token_ids.tolist()
|
||||
accept_lens = result.accept_lens.tolist()
|
||||
result.num_accepted_tokens = sum(accept_lens) - len(batch.reqs)
|
||||
result.accept_length_per_req_cpu = [x - 1 for x in accept_lens]
|
||||
|
||||
predict_tokens = []
|
||||
stride = self.draft_worker.speculative_num_draft_tokens
|
||||
@@ -359,6 +361,9 @@ class SchedulerOutputProcessorMixin:
|
||||
req.output_ids.extend(next_token_id)
|
||||
new_accepted_len = len(next_token_id)
|
||||
|
||||
# Update Mamba last track seqlen
|
||||
self._mamba_prefix_cache_update(req, batch, result, i)
|
||||
|
||||
req.check_finished(new_accepted_len)
|
||||
|
||||
if req.finished():
|
||||
@@ -424,6 +429,31 @@ class SchedulerOutputProcessorMixin:
|
||||
):
|
||||
self.log_decode_stats(can_run_cuda_graph, running_batch=batch)
|
||||
|
||||
def _mamba_prefix_cache_update(
|
||||
self, req: Req, batch: ScheduleBatch, result: GenerationBatchResult, i: int
|
||||
) -> None:
|
||||
seq_len = len(req.origin_input_ids) + len(req.output_ids) - 1
|
||||
if req.mamba_ping_pong_track_buffer is not None:
|
||||
mamba_track_interval = get_global_server_args().mamba_track_interval
|
||||
if batch.spec_algorithm.is_none() and seq_len % mamba_track_interval == 0:
|
||||
# for non-spec decode, we update mamba_last_track_seqlen at the end of each track interval
|
||||
req.mamba_next_track_idx = 1 - req.mamba_next_track_idx
|
||||
req.mamba_last_track_seqlen = seq_len
|
||||
elif (
|
||||
not batch.spec_algorithm.is_none()
|
||||
and result.accept_length_per_req_cpu is not None
|
||||
):
|
||||
# for spec decode, update mamba_last_track_seqlen if this iteration crosses a track interval
|
||||
actual_seq_len = req.seqlen - 1
|
||||
if (
|
||||
actual_seq_len // mamba_track_interval
|
||||
!= (actual_seq_len - result.accept_length_per_req_cpu[i])
|
||||
// mamba_track_interval
|
||||
):
|
||||
req.mamba_last_track_seqlen = (
|
||||
actual_seq_len // mamba_track_interval * mamba_track_interval
|
||||
)
|
||||
|
||||
def _process_input_token_logprobs(
|
||||
self, req: Req, input_token_logprobs: List
|
||||
) -> None:
|
||||
|
||||
@@ -121,9 +121,22 @@ class SchedulerRuntimeCheckerMixin:
|
||||
full_num_used != self.tree_cache.full_protected_size()
|
||||
or mamba_num_used != self.tree_cache.mamba_protected_size()
|
||||
)
|
||||
free_full_pages = set(
|
||||
self.token_to_kv_pool_allocator.free_pages.tolist()
|
||||
+ self.token_to_kv_pool_allocator.release_pages.tolist()
|
||||
)
|
||||
cached_full_pages = set(self.tree_cache.all_values_flatten().tolist())
|
||||
expected_full_pages = set(range(1, self.token_to_kv_pool_allocator.size + 1))
|
||||
leaked_full_pages = expected_full_pages - free_full_pages - cached_full_pages
|
||||
free_mamba_pages = set(self.req_to_token_pool.mamba_pool.free_slots.tolist())
|
||||
cached_mamba_pages = set(self.tree_cache.all_mamba_values_flatten().tolist())
|
||||
expected_mamba_pages = set(range(self.req_to_token_pool.mamba_pool.size))
|
||||
leaked_mamba_pages = (
|
||||
expected_mamba_pages - free_mamba_pages - cached_mamba_pages
|
||||
)
|
||||
token_msg = (
|
||||
f"{full_available_size=}, {full_evictable_size=}, {self.token_to_kv_pool_allocator.size=}, {self.tree_cache.full_protected_size()=}\n"
|
||||
f"{mamba_available_size=}, {mamba_evictable_size=}, {self.req_to_token_pool.mamba_pool.size=}, {self.tree_cache.mamba_protected_size()=}\n"
|
||||
f"{mamba_available_size=}, {mamba_evictable_size=}, {self.req_to_token_pool.mamba_pool.size=}, {self.tree_cache.mamba_protected_size()=}, leaked_full_pages={leaked_full_pages if len(leaked_full_pages) > 0 else None}, leaked_mamba_pages={leaked_mamba_pages if len(leaked_mamba_pages) > 0 else None}\n"
|
||||
)
|
||||
return memory_leak, token_msg
|
||||
|
||||
@@ -207,7 +220,7 @@ class SchedulerRuntimeCheckerMixin:
|
||||
def check_memory(self: Scheduler):
|
||||
if self.is_hybrid_swa:
|
||||
memory_leak, token_msg = self._check_hybrid_memory()
|
||||
elif self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache):
|
||||
elif self.is_hybrid_ssm and isinstance(self.tree_cache, MambaRadixCache):
|
||||
memory_leak, token_msg = self._check_mamba_memory()
|
||||
else:
|
||||
memory_leak, token_msg = self._check_radix_cache_memory()
|
||||
@@ -242,7 +255,7 @@ class SchedulerRuntimeCheckerMixin:
|
||||
) = self._get_swa_token_info()
|
||||
num_used = max(full_num_used, swa_num_used)
|
||||
token_usage = max(full_token_usage, swa_token_usage)
|
||||
elif self.is_ssm_model:
|
||||
elif self.is_hybrid_ssm:
|
||||
(
|
||||
num_used,
|
||||
_,
|
||||
@@ -281,7 +294,7 @@ class SchedulerRuntimeCheckerMixin:
|
||||
|
||||
def check_tree_cache(self: Scheduler):
|
||||
if (self.is_hybrid_swa and isinstance(self.tree_cache, SWARadixCache)) or (
|
||||
self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache)
|
||||
self.is_hybrid_ssm and isinstance(self.tree_cache, MambaRadixCache)
|
||||
):
|
||||
self.tree_cache.sanity_check()
|
||||
|
||||
@@ -344,7 +357,7 @@ class SchedulerWatchdog:
|
||||
# Print batch size and memory pool info to check whether there are de-sync issues.
|
||||
if self.scheduler.is_hybrid_swa:
|
||||
_, info_msg = self.scheduler._check_hybrid_memory()
|
||||
elif self.scheduler.is_ssm_model and isinstance(
|
||||
elif self.scheduler.is_hybrid_ssm and isinstance(
|
||||
self.scheduler.tree_cache, MambaRadixCache
|
||||
):
|
||||
_, info_msg = self.scheduler._check_mamba_memory()
|
||||
|
||||
@@ -24,7 +24,8 @@ class GenerationBatchResult:
|
||||
logits_output: Optional[LogitsProcessorOutput] = None
|
||||
pp_hidden_states_proxy_tensors: Optional[PPProxyTensors] = None
|
||||
next_token_ids: Optional[torch.Tensor] = None
|
||||
num_accepted_tokens: Optional[int] = None
|
||||
num_accepted_tokens: int = 0
|
||||
accept_length_per_req_cpu: Optional[List[int]] = None
|
||||
can_run_cuda_graph: bool = False
|
||||
|
||||
# For output processing
|
||||
|
||||
Reference in New Issue
Block a user