[Auto Sync] Rename is_hybrid to is_hybrid_swa (#14252)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com> Co-authored-by: Hanming Lu <69857889+hanming-lu@users.noreply.github.com> Co-authored-by: Hanming Lu <hanming@x.ai>
This commit is contained in:
co-authored by
github-actions[bot] <github-actions[bot]@users.noreply.github.com>
Hanming Lu
Hanming Lu
parent
63b9300f00
commit
64092c8b55
@@ -1066,7 +1066,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
req_to_token_pool: ReqToTokenPool = None
|
||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator = None
|
||||
tree_cache: BasePrefixCache = None
|
||||
is_hybrid: bool = False
|
||||
is_hybrid_swa: bool = False
|
||||
|
||||
# Batch configs
|
||||
model_config: ModelConfig = None
|
||||
@@ -1189,21 +1189,21 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
):
|
||||
return_logprob = any(req.return_logprob for req in reqs)
|
||||
|
||||
is_hybrid = False
|
||||
is_hybrid_swa = False
|
||||
if isinstance(token_to_kv_pool_allocator, SWATokenToKVPoolAllocator):
|
||||
assert (
|
||||
tree_cache is None
|
||||
or isinstance(tree_cache, SWARadixCache)
|
||||
or isinstance(tree_cache, SWAChunkCache)
|
||||
), "SWARadixCache or SWAChunkCache is required for SWATokenToKVPoolAllocator"
|
||||
is_hybrid = True
|
||||
is_hybrid_swa = True
|
||||
|
||||
return cls(
|
||||
reqs=reqs,
|
||||
req_to_token_pool=req_to_token_pool,
|
||||
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
||||
tree_cache=tree_cache,
|
||||
is_hybrid=is_hybrid,
|
||||
is_hybrid_swa=is_hybrid_swa,
|
||||
model_config=model_config,
|
||||
enable_overlap=enable_overlap,
|
||||
return_logprob=return_logprob,
|
||||
@@ -1612,7 +1612,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
):
|
||||
if len(sorted_indices) == 1:
|
||||
# Corner case: only one request left
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
full_available_size = (
|
||||
self.token_to_kv_pool_allocator.full_available_size()
|
||||
)
|
||||
@@ -1978,7 +1978,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
)
|
||||
|
||||
def _is_available_size_sufficient(self, num_tokens: int) -> bool:
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
return (
|
||||
self.token_to_kv_pool_allocator.full_available_size() >= num_tokens
|
||||
and self.token_to_kv_pool_allocator.swa_available_size() >= num_tokens
|
||||
|
||||
@@ -359,7 +359,7 @@ class PrefillAdder:
|
||||
]
|
||||
)
|
||||
|
||||
self.is_hybrid = isinstance(
|
||||
self.is_hybrid_swa = isinstance(
|
||||
self.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator
|
||||
)
|
||||
self.is_hybrid_gdn_cache = isinstance(self.tree_cache, MambaRadixCache)
|
||||
@@ -380,7 +380,7 @@ class PrefillAdder:
|
||||
|
||||
@property
|
||||
def rem_total_tokens(self):
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
available_and_evictable = min(
|
||||
self.token_to_kv_pool_allocator.full_available_size()
|
||||
+ self.tree_cache.full_evictable_size(),
|
||||
@@ -402,7 +402,7 @@ class PrefillAdder:
|
||||
|
||||
@property
|
||||
def cur_rem_tokens(self):
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
available_and_evictable = min(
|
||||
self.token_to_kv_pool_allocator.full_available_size()
|
||||
+ self.tree_cache.full_evictable_size(),
|
||||
@@ -472,7 +472,7 @@ class PrefillAdder:
|
||||
|
||||
@contextmanager
|
||||
def _lock_node(self, last_node: TreeNode):
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
try:
|
||||
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(last_node)
|
||||
yield None
|
||||
@@ -525,7 +525,7 @@ class PrefillAdder:
|
||||
else:
|
||||
add_req_state(req, insert_sort=True)
|
||||
|
||||
if not self.is_hybrid:
|
||||
if not self.is_hybrid_swa:
|
||||
# Skip this logic for swa. The SWA has different memory management, and
|
||||
# this mechanism is underestimating the memory usage.
|
||||
cur_rem_tokens = self.cur_rem_tokens - self.ceil_paged_tokens(
|
||||
@@ -616,7 +616,7 @@ class PrefillAdder:
|
||||
if self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
|
||||
# Non-chunked prefill
|
||||
self.can_run_list.append(req)
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(req.last_node)
|
||||
req.swa_uuid_for_lock = swa_uuid_for_lock
|
||||
else:
|
||||
@@ -652,7 +652,7 @@ class PrefillAdder:
|
||||
|
||||
self.can_run_list.append(req)
|
||||
self.new_chunked_req = req
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(req.last_node)
|
||||
req.swa_uuid_for_lock = swa_uuid_for_lock
|
||||
else:
|
||||
|
||||
@@ -397,10 +397,10 @@ class Scheduler(
|
||||
set_random_seed(self.random_seed)
|
||||
|
||||
# Hybrid memory pool
|
||||
self.is_hybrid = self.tp_worker.is_hybrid
|
||||
self.is_hybrid_swa = self.tp_worker.is_hybrid_swa
|
||||
self.is_hybrid_gdn = self.tp_worker.model_runner.hybrid_gdn_config is not None
|
||||
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
self.sliding_window_size = self.tp_worker.sliding_window_size
|
||||
self.full_tokens_per_layer, self.swa_tokens_per_layer = (
|
||||
self.tp_worker.get_tokens_per_layer_info()
|
||||
@@ -732,7 +732,7 @@ class Scheduler(
|
||||
server_args.chunked_prefill_size is not None
|
||||
and server_args.disable_radix_cache
|
||||
):
|
||||
if not self.is_hybrid:
|
||||
if not self.is_hybrid_swa:
|
||||
from sglang.srt.mem_cache.chunk_cache import ChunkCache
|
||||
|
||||
self.tree_cache = ChunkCache(params)
|
||||
@@ -756,7 +756,7 @@ class Scheduler(
|
||||
self.tp_worker.register_hicache_layer_transfer_counter(
|
||||
self.tree_cache.cache_controller.layer_done_counter
|
||||
)
|
||||
elif self.is_hybrid:
|
||||
elif self.is_hybrid_swa:
|
||||
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
|
||||
|
||||
self.tree_cache = SWARadixCache(
|
||||
|
||||
@@ -95,7 +95,7 @@ class SchedulerMetricsMixin:
|
||||
self.last_prefill_tokens = adder.log_input_tokens
|
||||
|
||||
# TODO: generalize this for various memory pools
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
(
|
||||
full_num_used,
|
||||
swa_num_used,
|
||||
@@ -164,7 +164,7 @@ class SchedulerMetricsMixin:
|
||||
self.stats.num_running_reqs_offline_batch = running_bs_offline_batch
|
||||
self.stats.num_used_tokens = num_used
|
||||
self.stats.token_usage = token_usage
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
self.stats.swa_token_usage = swa_token_usage
|
||||
if self.is_hybrid_gdn:
|
||||
self.stats.mamba_usage = mamba_usage
|
||||
@@ -219,7 +219,7 @@ class SchedulerMetricsMixin:
|
||||
num_running_reqs_offline_batch = 0
|
||||
|
||||
# TODO: generalize this for various memory pools
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
(
|
||||
full_num_used,
|
||||
swa_num_used,
|
||||
@@ -313,7 +313,7 @@ class SchedulerMetricsMixin:
|
||||
self.stats.num_running_reqs_offline_batch = num_running_reqs_offline_batch
|
||||
self.stats.num_used_tokens = num_used
|
||||
self.stats.token_usage = token_usage
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
self.stats.swa_token_usage = swa_token_usage
|
||||
if self.is_hybrid_gdn:
|
||||
self.stats.mamba_usage = mamba_usage
|
||||
@@ -398,7 +398,7 @@ class SchedulerMetricsMixin:
|
||||
)
|
||||
|
||||
def get_load(self: Scheduler, _: GetLoadReqInput = None) -> GetLoadReqOutput:
|
||||
if self.is_hybrid:
|
||||
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_hybrid_gdn:
|
||||
|
||||
@@ -202,7 +202,7 @@ class SchedulerRuntimeCheckerMixin:
|
||||
)
|
||||
|
||||
def check_memory(self: Scheduler):
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
memory_leak, token_msg = self._check_hybrid_memory()
|
||||
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
||||
memory_leak, token_msg = self._check_mamba_memory()
|
||||
@@ -226,7 +226,7 @@ class SchedulerRuntimeCheckerMixin:
|
||||
and time.perf_counter() > self.metrics_collector.last_log_time + 30
|
||||
):
|
||||
# During idle time, also collect metrics every 30 seconds.
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
(
|
||||
full_num_used,
|
||||
swa_num_used,
|
||||
@@ -277,7 +277,7 @@ class SchedulerRuntimeCheckerMixin:
|
||||
self._publish_kv_events()
|
||||
|
||||
def check_tree_cache(self: Scheduler):
|
||||
if (self.is_hybrid and isinstance(self.tree_cache, SWARadixCache)) or (
|
||||
if (self.is_hybrid_swa and isinstance(self.tree_cache, SWARadixCache)) or (
|
||||
self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache)
|
||||
):
|
||||
self.tree_cache.sanity_check()
|
||||
@@ -320,7 +320,7 @@ class SchedulerRuntimeCheckerMixin:
|
||||
|
||||
if not disable_request_logging():
|
||||
# Print batch size and memory pool info to check whether there are de-sync issues.
|
||||
if self.is_hybrid:
|
||||
if self.is_hybrid_swa:
|
||||
_, info_msg = self._check_hybrid_memory()
|
||||
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
||||
_, info_msg = self._check_mamba_memory()
|
||||
|
||||
@@ -72,8 +72,8 @@ class BaseTpWorker(ABC):
|
||||
return self.model_runner.sliding_window_size
|
||||
|
||||
@property
|
||||
def is_hybrid(self) -> bool:
|
||||
return self.model_runner.is_hybrid is not None
|
||||
def is_hybrid_swa(self) -> bool:
|
||||
return self.model_runner.is_hybrid_swa is not None
|
||||
|
||||
def get_tokens_per_layer_info(self):
|
||||
return (
|
||||
|
||||
Reference in New Issue
Block a user