[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]
Hanming Lu
Hanming Lu
parent
63b9300f00
commit
64092c8b55
@@ -154,13 +154,13 @@ class ModelConfig:
|
|||||||
self.attention_chunk_size = getattr(
|
self.attention_chunk_size = getattr(
|
||||||
self.hf_text_config, "attention_chunk_size", None
|
self.hf_text_config, "attention_chunk_size", None
|
||||||
)
|
)
|
||||||
self.is_hybrid = is_hybrid_model(
|
self.is_hybrid_swa = is_hybrid_model(
|
||||||
self.hf_config.architectures,
|
self.hf_config.architectures,
|
||||||
hybrid_kvcache_ratio=hybrid_kvcache_ratio,
|
hybrid_kvcache_ratio=hybrid_kvcache_ratio,
|
||||||
context_length=context_length,
|
context_length=context_length,
|
||||||
attention_chunk_size=self.attention_chunk_size,
|
attention_chunk_size=self.attention_chunk_size,
|
||||||
)
|
)
|
||||||
if self.is_hybrid is not None:
|
if self.is_hybrid_swa is not None:
|
||||||
self.swa_attention_layer_ids, self.full_attention_layer_ids = (
|
self.swa_attention_layer_ids, self.full_attention_layer_ids = (
|
||||||
get_hybrid_layer_ids(
|
get_hybrid_layer_ids(
|
||||||
self.hf_config.architectures, self.hf_text_config.num_hidden_layers
|
self.hf_config.architectures, self.hf_text_config.num_hidden_layers
|
||||||
|
|||||||
@@ -569,7 +569,7 @@ class DecodePreallocQueue:
|
|||||||
else 0
|
else 0
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.scheduler.model_config.is_hybrid:
|
if self.scheduler.model_config.is_hybrid_swa:
|
||||||
available_size = min(
|
available_size = min(
|
||||||
self.token_to_kv_pool_allocator.full_available_size(),
|
self.token_to_kv_pool_allocator.full_available_size(),
|
||||||
self.token_to_kv_pool_allocator.swa_available_size(),
|
self.token_to_kv_pool_allocator.swa_available_size(),
|
||||||
|
|||||||
@@ -328,8 +328,8 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||||
self.skip_prefill = skip_prefill
|
self.skip_prefill = skip_prefill
|
||||||
self.is_hybrid = model_runner.is_hybrid
|
self.is_hybrid_swa = model_runner.is_hybrid_swa
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
self.full_to_swa_index_mapping = (
|
self.full_to_swa_index_mapping = (
|
||||||
model_runner.token_to_kv_pool.full_to_swa_index_mapping
|
model_runner.token_to_kv_pool.full_to_swa_index_mapping
|
||||||
)
|
)
|
||||||
@@ -720,10 +720,10 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
# Calculate window size (can be moved to metadata if layer properties don't change)
|
# Calculate window size (can be moved to metadata if layer properties don't change)
|
||||||
# we don't do layer.sliding_window_size - 1 since in model.get_attention_sliding_window_size() we already - 1
|
# we don't do layer.sliding_window_size - 1 since in model.get_attention_sliding_window_size() we already - 1
|
||||||
# here is two side inclusive
|
# here is two side inclusive
|
||||||
is_swa = (
|
is_hybrid_swa = (
|
||||||
layer.sliding_window_size is not None and layer.sliding_window_size > -1
|
layer.sliding_window_size is not None and layer.sliding_window_size > -1
|
||||||
)
|
)
|
||||||
window_size = (layer.sliding_window_size, 0) if is_swa else (-1, -1)
|
window_size = (layer.sliding_window_size, 0) if is_hybrid_swa else (-1, -1)
|
||||||
k_descale, v_descale = None, None
|
k_descale, v_descale = None, None
|
||||||
# only use kv scaling if: 1) fp8 kv is explicitly enabled, 2) RadixAttention
|
# only use kv scaling if: 1) fp8 kv is explicitly enabled, 2) RadixAttention
|
||||||
# has corresponding quantization method so that layer.k_scale is not None,
|
# has corresponding quantization method so that layer.k_scale is not None,
|
||||||
@@ -759,7 +759,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
use_cascade_attn = (
|
use_cascade_attn = (
|
||||||
forward_batch.forward_mode.is_target_verify()
|
forward_batch.forward_mode.is_target_verify()
|
||||||
and self.topk > 1
|
and self.topk > 1
|
||||||
and not is_swa
|
and not is_hybrid_swa
|
||||||
)
|
)
|
||||||
|
|
||||||
# For fa3 interface version compatibility, we put new fields into conditional keyword args
|
# For fa3 interface version compatibility, we put new fields into conditional keyword args
|
||||||
@@ -776,7 +776,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
cu_seqlens_q = local_metadata.local_query_start_loc
|
cu_seqlens_q = local_metadata.local_query_start_loc
|
||||||
cache_seqlens = local_metadata.local_seqused_k
|
cache_seqlens = local_metadata.local_seqused_k
|
||||||
max_seqlen_q = local_metadata.local_max_query_len
|
max_seqlen_q = local_metadata.local_max_query_len
|
||||||
elif is_swa and metadata.swa_spec_metadata is not None:
|
elif is_hybrid_swa and metadata.swa_spec_metadata is not None:
|
||||||
swa_spec_metadata = metadata.swa_spec_metadata
|
swa_spec_metadata = metadata.swa_spec_metadata
|
||||||
page_table = swa_spec_metadata.page_table
|
page_table = swa_spec_metadata.page_table
|
||||||
cu_seqlens_q = swa_spec_metadata.cu_seqlens_q
|
cu_seqlens_q = swa_spec_metadata.cu_seqlens_q
|
||||||
@@ -2079,7 +2079,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
|
|
||||||
cu_seqlens_q = metadata.cu_seqlens_q
|
cu_seqlens_q = metadata.cu_seqlens_q
|
||||||
cache_seqlens_int32 = metadata.cache_seqlens_int32
|
cache_seqlens_int32 = metadata.cache_seqlens_int32
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
page_table = self.full_to_swa_index_mapping[metadata.page_table].to(
|
page_table = self.full_to_swa_index_mapping[metadata.page_table].to(
|
||||||
torch.int32
|
torch.int32
|
||||||
)
|
)
|
||||||
@@ -2203,7 +2203,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
# Without this slicing, the pre-allocated page_table may contain zeros or invalid indices
|
# Without this slicing, the pre-allocated page_table may contain zeros or invalid indices
|
||||||
# beyond the actual sequence length, leading to incorrect attention calculations
|
# beyond the actual sequence length, leading to incorrect attention calculations
|
||||||
max_seq_len = int(seqlens.max().item())
|
max_seq_len = int(seqlens.max().item())
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
sliced_page_table = self.full_to_swa_index_mapping[
|
sliced_page_table = self.full_to_swa_index_mapping[
|
||||||
metadata.page_table[:bs, :max_seq_len]
|
metadata.page_table[:bs, :max_seq_len]
|
||||||
].to(torch.int32)
|
].to(torch.int32)
|
||||||
|
|||||||
@@ -64,8 +64,8 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
self.use_mla is False
|
self.use_mla is False
|
||||||
), "XPUAttentionBackend doesn't support MLA yet, please use --attention-backend triton instead."
|
), "XPUAttentionBackend doesn't support MLA yet, please use --attention-backend triton instead."
|
||||||
self.skip_prefill = skip_prefill
|
self.skip_prefill = skip_prefill
|
||||||
self.is_hybrid = model_runner.is_hybrid
|
self.is_hybrid_swa = model_runner.is_hybrid_swa
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
self.full_to_swa_index_mapping = (
|
self.full_to_swa_index_mapping = (
|
||||||
model_runner.token_to_kv_pool.full_to_swa_index_mapping
|
model_runner.token_to_kv_pool.full_to_swa_index_mapping
|
||||||
)
|
)
|
||||||
@@ -416,10 +416,10 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
# Calculate window size (can be moved to metadata if layer properties don't change)
|
# Calculate window size (can be moved to metadata if layer properties don't change)
|
||||||
# we don't do layer.sliding_window_size - 1 since in model.get_attention_sliding_window_size() we already - 1
|
# we don't do layer.sliding_window_size - 1 since in model.get_attention_sliding_window_size() we already - 1
|
||||||
# here is two side inclusive
|
# here is two side inclusive
|
||||||
is_swa = (
|
is_hybrid_swa = (
|
||||||
layer.sliding_window_size is not None and layer.sliding_window_size > -1
|
layer.sliding_window_size is not None and layer.sliding_window_size > -1
|
||||||
)
|
)
|
||||||
window_size = (layer.sliding_window_size, 0) if is_swa else (-1, -1)
|
window_size = (layer.sliding_window_size, 0) if is_hybrid_swa else (-1, -1)
|
||||||
|
|
||||||
# currently no FP8 KV cache supported
|
# currently no FP8 KV cache supported
|
||||||
k_descale, v_descale = None, None
|
k_descale, v_descale = None, None
|
||||||
@@ -450,7 +450,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
use_cascade_attn = (
|
use_cascade_attn = (
|
||||||
forward_batch.forward_mode.is_target_verify()
|
forward_batch.forward_mode.is_target_verify()
|
||||||
and self.topk > 1
|
and self.topk > 1
|
||||||
and not is_swa
|
and not is_hybrid_swa
|
||||||
)
|
)
|
||||||
|
|
||||||
# For fa3 interface version compatibility, we put new fields into conditional keyword args
|
# For fa3 interface version compatibility, we put new fields into conditional keyword args
|
||||||
@@ -465,7 +465,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
cu_seqlens_q = local_metadata.local_query_start_loc
|
cu_seqlens_q = local_metadata.local_query_start_loc
|
||||||
cache_seqlens = local_metadata.local_seqused_k
|
cache_seqlens = local_metadata.local_seqused_k
|
||||||
max_seqlen_q = local_metadata.local_max_query_len
|
max_seqlen_q = local_metadata.local_max_query_len
|
||||||
elif is_swa and metadata.swa_spec_metadata is not None:
|
elif is_hybrid_swa and metadata.swa_spec_metadata is not None:
|
||||||
swa_spec_metadata = metadata.swa_spec_metadata
|
swa_spec_metadata = metadata.swa_spec_metadata
|
||||||
page_table = swa_spec_metadata.page_table
|
page_table = swa_spec_metadata.page_table
|
||||||
cu_seqlens_q = swa_spec_metadata.cu_seqlens_q
|
cu_seqlens_q = swa_spec_metadata.cu_seqlens_q
|
||||||
@@ -942,7 +942,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
|
|
||||||
cu_seqlens_q = metadata.cu_seqlens_q
|
cu_seqlens_q = metadata.cu_seqlens_q
|
||||||
cache_seqlens_int32 = metadata.cache_seqlens_int32
|
cache_seqlens_int32 = metadata.cache_seqlens_int32
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
page_table = self.full_to_swa_index_mapping[metadata.page_table].to(
|
page_table = self.full_to_swa_index_mapping[metadata.page_table].to(
|
||||||
torch.int32
|
torch.int32
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1066,7 +1066,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
req_to_token_pool: ReqToTokenPool = None
|
req_to_token_pool: ReqToTokenPool = None
|
||||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator = None
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator = None
|
||||||
tree_cache: BasePrefixCache = None
|
tree_cache: BasePrefixCache = None
|
||||||
is_hybrid: bool = False
|
is_hybrid_swa: bool = False
|
||||||
|
|
||||||
# Batch configs
|
# Batch configs
|
||||||
model_config: ModelConfig = None
|
model_config: ModelConfig = None
|
||||||
@@ -1189,21 +1189,21 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
):
|
):
|
||||||
return_logprob = any(req.return_logprob for req in reqs)
|
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):
|
if isinstance(token_to_kv_pool_allocator, SWATokenToKVPoolAllocator):
|
||||||
assert (
|
assert (
|
||||||
tree_cache is None
|
tree_cache is None
|
||||||
or isinstance(tree_cache, SWARadixCache)
|
or isinstance(tree_cache, SWARadixCache)
|
||||||
or isinstance(tree_cache, SWAChunkCache)
|
or isinstance(tree_cache, SWAChunkCache)
|
||||||
), "SWARadixCache or SWAChunkCache is required for SWATokenToKVPoolAllocator"
|
), "SWARadixCache or SWAChunkCache is required for SWATokenToKVPoolAllocator"
|
||||||
is_hybrid = True
|
is_hybrid_swa = True
|
||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
reqs=reqs,
|
reqs=reqs,
|
||||||
req_to_token_pool=req_to_token_pool,
|
req_to_token_pool=req_to_token_pool,
|
||||||
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
||||||
tree_cache=tree_cache,
|
tree_cache=tree_cache,
|
||||||
is_hybrid=is_hybrid,
|
is_hybrid_swa=is_hybrid_swa,
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
enable_overlap=enable_overlap,
|
enable_overlap=enable_overlap,
|
||||||
return_logprob=return_logprob,
|
return_logprob=return_logprob,
|
||||||
@@ -1612,7 +1612,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
):
|
):
|
||||||
if len(sorted_indices) == 1:
|
if len(sorted_indices) == 1:
|
||||||
# Corner case: only one request left
|
# Corner case: only one request left
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
full_available_size = (
|
full_available_size = (
|
||||||
self.token_to_kv_pool_allocator.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:
|
def _is_available_size_sufficient(self, num_tokens: int) -> bool:
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
return (
|
return (
|
||||||
self.token_to_kv_pool_allocator.full_available_size() >= num_tokens
|
self.token_to_kv_pool_allocator.full_available_size() >= num_tokens
|
||||||
and self.token_to_kv_pool_allocator.swa_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.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator
|
||||||
)
|
)
|
||||||
self.is_hybrid_gdn_cache = isinstance(self.tree_cache, MambaRadixCache)
|
self.is_hybrid_gdn_cache = isinstance(self.tree_cache, MambaRadixCache)
|
||||||
@@ -380,7 +380,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def rem_total_tokens(self):
|
def rem_total_tokens(self):
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
available_and_evictable = min(
|
available_and_evictable = min(
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
self.token_to_kv_pool_allocator.full_available_size()
|
||||||
+ self.tree_cache.full_evictable_size(),
|
+ self.tree_cache.full_evictable_size(),
|
||||||
@@ -402,7 +402,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def cur_rem_tokens(self):
|
def cur_rem_tokens(self):
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
available_and_evictable = min(
|
available_and_evictable = min(
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
self.token_to_kv_pool_allocator.full_available_size()
|
||||||
+ self.tree_cache.full_evictable_size(),
|
+ self.tree_cache.full_evictable_size(),
|
||||||
@@ -472,7 +472,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def _lock_node(self, last_node: TreeNode):
|
def _lock_node(self, last_node: TreeNode):
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
try:
|
try:
|
||||||
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(last_node)
|
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(last_node)
|
||||||
yield None
|
yield None
|
||||||
@@ -525,7 +525,7 @@ class PrefillAdder:
|
|||||||
else:
|
else:
|
||||||
add_req_state(req, insert_sort=True)
|
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
|
# Skip this logic for swa. The SWA has different memory management, and
|
||||||
# this mechanism is underestimating the memory usage.
|
# this mechanism is underestimating the memory usage.
|
||||||
cur_rem_tokens = self.cur_rem_tokens - self.ceil_paged_tokens(
|
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:
|
if self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
|
||||||
# Non-chunked prefill
|
# Non-chunked prefill
|
||||||
self.can_run_list.append(req)
|
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)
|
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(req.last_node)
|
||||||
req.swa_uuid_for_lock = swa_uuid_for_lock
|
req.swa_uuid_for_lock = swa_uuid_for_lock
|
||||||
else:
|
else:
|
||||||
@@ -652,7 +652,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
self.can_run_list.append(req)
|
self.can_run_list.append(req)
|
||||||
self.new_chunked_req = 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)
|
swa_uuid_for_lock = self.tree_cache.inc_lock_ref(req.last_node)
|
||||||
req.swa_uuid_for_lock = swa_uuid_for_lock
|
req.swa_uuid_for_lock = swa_uuid_for_lock
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -397,10 +397,10 @@ class Scheduler(
|
|||||||
set_random_seed(self.random_seed)
|
set_random_seed(self.random_seed)
|
||||||
|
|
||||||
# Hybrid memory pool
|
# 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
|
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.sliding_window_size = self.tp_worker.sliding_window_size
|
||||||
self.full_tokens_per_layer, self.swa_tokens_per_layer = (
|
self.full_tokens_per_layer, self.swa_tokens_per_layer = (
|
||||||
self.tp_worker.get_tokens_per_layer_info()
|
self.tp_worker.get_tokens_per_layer_info()
|
||||||
@@ -732,7 +732,7 @@ class Scheduler(
|
|||||||
server_args.chunked_prefill_size is not None
|
server_args.chunked_prefill_size is not None
|
||||||
and server_args.disable_radix_cache
|
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
|
from sglang.srt.mem_cache.chunk_cache import ChunkCache
|
||||||
|
|
||||||
self.tree_cache = ChunkCache(params)
|
self.tree_cache = ChunkCache(params)
|
||||||
@@ -756,7 +756,7 @@ class Scheduler(
|
|||||||
self.tp_worker.register_hicache_layer_transfer_counter(
|
self.tp_worker.register_hicache_layer_transfer_counter(
|
||||||
self.tree_cache.cache_controller.layer_done_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
|
from sglang.srt.mem_cache.swa_radix_cache import SWARadixCache
|
||||||
|
|
||||||
self.tree_cache = SWARadixCache(
|
self.tree_cache = SWARadixCache(
|
||||||
|
|||||||
@@ -95,7 +95,7 @@ class SchedulerMetricsMixin:
|
|||||||
self.last_prefill_tokens = adder.log_input_tokens
|
self.last_prefill_tokens = adder.log_input_tokens
|
||||||
|
|
||||||
# TODO: generalize this for various memory pools
|
# TODO: generalize this for various memory pools
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
(
|
(
|
||||||
full_num_used,
|
full_num_used,
|
||||||
swa_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_running_reqs_offline_batch = running_bs_offline_batch
|
||||||
self.stats.num_used_tokens = num_used
|
self.stats.num_used_tokens = num_used
|
||||||
self.stats.token_usage = token_usage
|
self.stats.token_usage = token_usage
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
self.stats.swa_token_usage = swa_token_usage
|
self.stats.swa_token_usage = swa_token_usage
|
||||||
if self.is_hybrid_gdn:
|
if self.is_hybrid_gdn:
|
||||||
self.stats.mamba_usage = mamba_usage
|
self.stats.mamba_usage = mamba_usage
|
||||||
@@ -219,7 +219,7 @@ class SchedulerMetricsMixin:
|
|||||||
num_running_reqs_offline_batch = 0
|
num_running_reqs_offline_batch = 0
|
||||||
|
|
||||||
# TODO: generalize this for various memory pools
|
# TODO: generalize this for various memory pools
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
(
|
(
|
||||||
full_num_used,
|
full_num_used,
|
||||||
swa_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_running_reqs_offline_batch = num_running_reqs_offline_batch
|
||||||
self.stats.num_used_tokens = num_used
|
self.stats.num_used_tokens = num_used
|
||||||
self.stats.token_usage = token_usage
|
self.stats.token_usage = token_usage
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
self.stats.swa_token_usage = swa_token_usage
|
self.stats.swa_token_usage = swa_token_usage
|
||||||
if self.is_hybrid_gdn:
|
if self.is_hybrid_gdn:
|
||||||
self.stats.mamba_usage = mamba_usage
|
self.stats.mamba_usage = mamba_usage
|
||||||
@@ -398,7 +398,7 @@ class SchedulerMetricsMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_load(self: Scheduler, _: GetLoadReqInput = None) -> GetLoadReqOutput:
|
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()
|
full_num_used, swa_num_used, *_ = self._get_swa_token_info()
|
||||||
num_tokens = max(full_num_used, swa_num_used)
|
num_tokens = max(full_num_used, swa_num_used)
|
||||||
elif self.is_hybrid_gdn:
|
elif self.is_hybrid_gdn:
|
||||||
|
|||||||
@@ -202,7 +202,7 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def check_memory(self: Scheduler):
|
def check_memory(self: Scheduler):
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
memory_leak, token_msg = self._check_hybrid_memory()
|
memory_leak, token_msg = self._check_hybrid_memory()
|
||||||
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
||||||
memory_leak, token_msg = self._check_mamba_memory()
|
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
|
and time.perf_counter() > self.metrics_collector.last_log_time + 30
|
||||||
):
|
):
|
||||||
# During idle time, also collect metrics every 30 seconds.
|
# During idle time, also collect metrics every 30 seconds.
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
(
|
(
|
||||||
full_num_used,
|
full_num_used,
|
||||||
swa_num_used,
|
swa_num_used,
|
||||||
@@ -277,7 +277,7 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
self._publish_kv_events()
|
self._publish_kv_events()
|
||||||
|
|
||||||
def check_tree_cache(self: Scheduler):
|
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.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache)
|
||||||
):
|
):
|
||||||
self.tree_cache.sanity_check()
|
self.tree_cache.sanity_check()
|
||||||
@@ -320,7 +320,7 @@ class SchedulerRuntimeCheckerMixin:
|
|||||||
|
|
||||||
if not disable_request_logging():
|
if not disable_request_logging():
|
||||||
# Print batch size and memory pool info to check whether there are de-sync issues.
|
# 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()
|
_, info_msg = self._check_hybrid_memory()
|
||||||
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
elif self.is_hybrid_gdn and isinstance(self.tree_cache, MambaRadixCache):
|
||||||
_, info_msg = self._check_mamba_memory()
|
_, info_msg = self._check_mamba_memory()
|
||||||
|
|||||||
@@ -72,8 +72,8 @@ class BaseTpWorker(ABC):
|
|||||||
return self.model_runner.sliding_window_size
|
return self.model_runner.sliding_window_size
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_hybrid(self) -> bool:
|
def is_hybrid_swa(self) -> bool:
|
||||||
return self.model_runner.is_hybrid is not None
|
return self.model_runner.is_hybrid_swa is not None
|
||||||
|
|
||||||
def get_tokens_per_layer_info(self):
|
def get_tokens_per_layer_info(self):
|
||||||
return (
|
return (
|
||||||
|
|||||||
@@ -1192,6 +1192,7 @@ class SWAKVPool(KVCache):
|
|||||||
layer_num=self.full_layer_nums,
|
layer_num=self.full_layer_nums,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
# {layer_id: (index, is_swa_layer)}
|
||||||
self.layers_mapping: Dict[int, Tuple[int, bool]] = {}
|
self.layers_mapping: Dict[int, Tuple[int, bool]] = {}
|
||||||
for full_attn_layer_id, global_layer_id in enumerate(full_attention_layer_ids):
|
for full_attn_layer_id, global_layer_id in enumerate(full_attention_layer_ids):
|
||||||
self.layers_mapping[global_layer_id] = (full_attn_layer_id, False)
|
self.layers_mapping[global_layer_id] = (full_attn_layer_id, False)
|
||||||
@@ -1229,22 +1230,22 @@ class SWAKVPool(KVCache):
|
|||||||
return swa_kv_data_ptrs, swa_kv_data_lens, swa_kv_item_lens
|
return swa_kv_data_ptrs, swa_kv_data_lens, swa_kv_item_lens
|
||||||
|
|
||||||
def get_key_buffer(self, layer_id: int):
|
def get_key_buffer(self, layer_id: int):
|
||||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
if is_swa:
|
if is_swa_layer:
|
||||||
return self.swa_kv_pool.get_key_buffer(layer_id_pool)
|
return self.swa_kv_pool.get_key_buffer(layer_id_pool)
|
||||||
else:
|
else:
|
||||||
return self.full_kv_pool.get_key_buffer(layer_id_pool)
|
return self.full_kv_pool.get_key_buffer(layer_id_pool)
|
||||||
|
|
||||||
def get_value_buffer(self, layer_id: int):
|
def get_value_buffer(self, layer_id: int):
|
||||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
if is_swa:
|
if is_swa_layer:
|
||||||
return self.swa_kv_pool.get_value_buffer(layer_id_pool)
|
return self.swa_kv_pool.get_value_buffer(layer_id_pool)
|
||||||
else:
|
else:
|
||||||
return self.full_kv_pool.get_value_buffer(layer_id_pool)
|
return self.full_kv_pool.get_value_buffer(layer_id_pool)
|
||||||
|
|
||||||
def get_kv_buffer(self, layer_id: int):
|
def get_kv_buffer(self, layer_id: int):
|
||||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
if is_swa:
|
if is_swa_layer:
|
||||||
return self.swa_kv_pool.get_kv_buffer(layer_id_pool)
|
return self.swa_kv_pool.get_kv_buffer(layer_id_pool)
|
||||||
else:
|
else:
|
||||||
return self.full_kv_pool.get_kv_buffer(layer_id_pool)
|
return self.full_kv_pool.get_kv_buffer(layer_id_pool)
|
||||||
@@ -1264,8 +1265,8 @@ class SWAKVPool(KVCache):
|
|||||||
):
|
):
|
||||||
|
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
layer_id_pool, is_swa = self.layers_mapping[layer_id]
|
layer_id_pool, is_swa_layer = self.layers_mapping[layer_id]
|
||||||
if is_swa:
|
if is_swa_layer:
|
||||||
if self.full_to_swa_index_mapping is not None:
|
if self.full_to_swa_index_mapping is not None:
|
||||||
loc = self.translate_loc_from_full_to_swa(loc)
|
loc = self.translate_loc_from_full_to_swa(loc)
|
||||||
self.swa_kv_pool.set_kv_buffer(
|
self.swa_kv_pool.set_kv_buffer(
|
||||||
|
|||||||
@@ -108,9 +108,9 @@ def get_last_access_time() -> float64:
|
|||||||
|
|
||||||
|
|
||||||
class LRUList:
|
class LRUList:
|
||||||
def __init__(self, swa: bool = False):
|
def __init__(self, is_swa_list: bool = False):
|
||||||
self.swa = swa
|
self.is_swa_list = is_swa_list
|
||||||
if self.swa:
|
if self.is_swa_list:
|
||||||
self.prv = "swa_prev"
|
self.prv = "swa_prev"
|
||||||
self.nxt = "swa_next"
|
self.nxt = "swa_next"
|
||||||
self.lock_ref = "swa_lock_ref"
|
self.lock_ref = "swa_lock_ref"
|
||||||
@@ -163,7 +163,7 @@ class LRUList:
|
|||||||
"""
|
"""
|
||||||
assert node.id in self.cache, f"Resetting node {node.id=} not in lru list"
|
assert node.id in self.cache, f"Resetting node {node.id=} not in lru list"
|
||||||
assert (
|
assert (
|
||||||
not self.swa or not node.swa_tombstone
|
not self.is_swa_list or not node.swa_tombstone
|
||||||
), f"Resetting swa tombstone node in swa lru list: {node.id=}"
|
), f"Resetting swa tombstone node in swa lru list: {node.id=}"
|
||||||
self._remove_node(node)
|
self._remove_node(node)
|
||||||
self._add_node(node)
|
self._add_node(node)
|
||||||
@@ -176,7 +176,7 @@ class LRUList:
|
|||||||
prev_node = self.head
|
prev_node = self.head
|
||||||
while node != root_node:
|
while node != root_node:
|
||||||
# for swa lru list, only reset non-tombstone nodes
|
# for swa lru list, only reset non-tombstone nodes
|
||||||
if not self.swa or not node.swa_tombstone:
|
if not self.is_swa_list or not node.swa_tombstone:
|
||||||
assert (
|
assert (
|
||||||
node.id in self.cache
|
node.id in self.cache
|
||||||
), f"Resetting node {node.id=} not in lru list when resetting node and parents mru"
|
), f"Resetting node {node.id=} not in lru list when resetting node and parents mru"
|
||||||
@@ -190,7 +190,7 @@ class LRUList:
|
|||||||
Insert a (new) node as most recently used
|
Insert a (new) node as most recently used
|
||||||
"""
|
"""
|
||||||
assert (
|
assert (
|
||||||
not self.swa or not node.swa_tombstone
|
not self.is_swa_list or not node.swa_tombstone
|
||||||
), f"Inserting swa tombstone node in swa lru list: {node.id=}"
|
), f"Inserting swa tombstone node in swa lru list: {node.id=}"
|
||||||
assert (
|
assert (
|
||||||
node.id not in self.cache
|
node.id not in self.cache
|
||||||
@@ -204,7 +204,7 @@ class LRUList:
|
|||||||
"""
|
"""
|
||||||
assert node.id in self.cache, f"Removing node {node.id=} not in lru list"
|
assert node.id in self.cache, f"Removing node {node.id=} not in lru list"
|
||||||
assert (
|
assert (
|
||||||
not self.swa or not node.swa_tombstone
|
not self.is_swa_list or not node.swa_tombstone
|
||||||
), f"Removing swa tombstone node from swa lru list: {node.id=}"
|
), f"Removing swa tombstone node from swa lru list: {node.id=}"
|
||||||
del self.cache[node.id]
|
del self.cache[node.id]
|
||||||
self._remove_node(node)
|
self._remove_node(node)
|
||||||
@@ -282,7 +282,7 @@ class LRUList:
|
|||||||
checking if the lru list is valid.
|
checking if the lru list is valid.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
if self.swa:
|
if self.is_swa_list:
|
||||||
nodes = tree_cache._collect_nontombstone_nodes()
|
nodes = tree_cache._collect_nontombstone_nodes()
|
||||||
else:
|
else:
|
||||||
nodes = tree_cache._collect_all_nodes()
|
nodes = tree_cache._collect_all_nodes()
|
||||||
@@ -303,7 +303,7 @@ class LRUList:
|
|||||||
continue
|
continue
|
||||||
assert (
|
assert (
|
||||||
x == x_lru
|
x == x_lru
|
||||||
), f"Incorrect LRU list, {self.swa=}, x: {x.id=} != x_lru: {x_lru.id=}"
|
), f"Incorrect LRU list, {self.is_swa_list=}, x: {x.id=} != x_lru: {x_lru.id=}"
|
||||||
assert (
|
assert (
|
||||||
x_lru.full_lock_ref == 0
|
x_lru.full_lock_ref == 0
|
||||||
), f"x_lru should not be locked when idle, {x_lru.full_lock_ref=}, {x_lru.swa_uuid=}, {x_lru.id=}"
|
), f"x_lru should not be locked when idle, {x_lru.full_lock_ref=}, {x_lru.swa_uuid=}, {x_lru.id=}"
|
||||||
@@ -312,7 +312,7 @@ class LRUList:
|
|||||||
), f"x_lru should not be locked when idle, {x_lru.swa_lock_ref=}, {x_lru.swa_uuid=}, {x_lru.id=}"
|
), f"x_lru should not be locked when idle, {x_lru.swa_lock_ref=}, {x_lru.swa_uuid=}, {x_lru.id=}"
|
||||||
x_lru = getattr(x, self.prv)
|
x_lru = getattr(x, self.prv)
|
||||||
|
|
||||||
if self.swa:
|
if self.is_swa_list:
|
||||||
evictable_size = tree_cache.swa_evictable_size()
|
evictable_size = tree_cache.swa_evictable_size()
|
||||||
lru_list_evictable_size = tree_cache.swa_lru_list_evictable_size()
|
lru_list_evictable_size = tree_cache.swa_lru_list_evictable_size()
|
||||||
else:
|
else:
|
||||||
@@ -321,7 +321,7 @@ class LRUList:
|
|||||||
|
|
||||||
assert (
|
assert (
|
||||||
evictable_size == lru_list_evictable_size
|
evictable_size == lru_list_evictable_size
|
||||||
), f"{self.swa=}, total nodes: {total_nodes}, total lru plus 1: {total_lru_plus_1}, evictable size: {evictable_size} != lru list evictable size: {lru_list_evictable_size}"
|
), f"{self.is_swa_list=}, total nodes: {total_nodes}, total lru plus 1: {total_lru_plus_1}, evictable size: {evictable_size} != lru list evictable size: {lru_list_evictable_size}"
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
msg = f"SWA Radix tree sanity check failed, ping @hanming-lu: {e}"
|
msg = f"SWA Radix tree sanity check failed, ping @hanming-lu: {e}"
|
||||||
logger.error(msg)
|
logger.error(msg)
|
||||||
@@ -373,8 +373,8 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
self.full_protected_size_ = 0
|
self.full_protected_size_ = 0
|
||||||
self.swa_protected_size_ = 0
|
self.swa_protected_size_ = 0
|
||||||
# LRU lists are used to maintain the order of eviction of the nodes in the tree
|
# LRU lists are used to maintain the order of eviction of the nodes in the tree
|
||||||
self.full_lru_list = LRUList(swa=False)
|
self.full_lru_list = LRUList(is_swa_list=False)
|
||||||
self.swa_lru_list = LRUList(swa=True)
|
self.swa_lru_list = LRUList(is_swa_list=True)
|
||||||
|
|
||||||
def match_prefix(self, key: RadixKey, **kwargs) -> MatchResult:
|
def match_prefix(self, key: RadixKey, **kwargs) -> MatchResult:
|
||||||
"""Find the matching prefix from the radix tree.
|
"""Find the matching prefix from the radix tree.
|
||||||
|
|||||||
@@ -315,8 +315,7 @@ class ModelRunner:
|
|||||||
self.page_size = server_args.page_size
|
self.page_size = server_args.page_size
|
||||||
self.req_to_token_pool = req_to_token_pool
|
self.req_to_token_pool = req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||||
self.is_hybrid = model_config.is_hybrid
|
self.is_hybrid_swa = model_config.is_hybrid_swa
|
||||||
self.is_hybrid_swa = self.is_hybrid
|
|
||||||
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
|
||||||
self.attention_chunk_size = model_config.attention_chunk_size
|
self.attention_chunk_size = model_config.attention_chunk_size
|
||||||
self.forward_pass_id = 0
|
self.forward_pass_id = 0
|
||||||
@@ -444,7 +443,7 @@ class ModelRunner:
|
|||||||
):
|
):
|
||||||
architectures = self.model_config.hf_config.architectures
|
architectures = self.model_config.hf_config.architectures
|
||||||
if architectures and not any("Llama4" in arch for arch in architectures):
|
if architectures and not any("Llama4" in arch for arch in architectures):
|
||||||
self.is_hybrid = self.model_config.is_hybrid = True
|
self.is_hybrid_swa = self.model_config.is_hybrid_swa = True
|
||||||
|
|
||||||
if config := self.mamba2_config:
|
if config := self.mamba2_config:
|
||||||
class_name = config.__class__.__name__
|
class_name = config.__class__.__name__
|
||||||
@@ -1530,8 +1529,8 @@ class ModelRunner:
|
|||||||
in self.model_config.hf_config.architectures
|
in self.model_config.hf_config.architectures
|
||||||
):
|
):
|
||||||
temp_ratio = (
|
temp_ratio = (
|
||||||
(1 - self.is_hybrid)
|
(1 - self.is_hybrid_swa)
|
||||||
+ self.is_hybrid
|
+ self.is_hybrid_swa
|
||||||
* self.attention_chunk_size
|
* self.attention_chunk_size
|
||||||
/ self.model_config.context_len
|
/ self.model_config.context_len
|
||||||
)
|
)
|
||||||
@@ -1567,7 +1566,7 @@ class ModelRunner:
|
|||||||
try:
|
try:
|
||||||
layers = self.model.language_model.layers
|
layers = self.model.language_model.layers
|
||||||
except:
|
except:
|
||||||
self.is_hybrid = False
|
self.is_hybrid_swa = False
|
||||||
return
|
return
|
||||||
|
|
||||||
for layer in layers:
|
for layer in layers:
|
||||||
@@ -1743,7 +1742,7 @@ class ModelRunner:
|
|||||||
self.max_total_num_tokens = tensor.item()
|
self.max_total_num_tokens = tensor.item()
|
||||||
|
|
||||||
# create token size for hybrid cache
|
# create token size for hybrid cache
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
self.set_num_token_hybrid()
|
self.set_num_token_hybrid()
|
||||||
|
|
||||||
if self.max_total_num_tokens <= 0:
|
if self.max_total_num_tokens <= 0:
|
||||||
@@ -1900,7 +1899,7 @@ class ModelRunner:
|
|||||||
end_layer=self.end_layer,
|
end_layer=self.end_layer,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
self.token_to_kv_pool = SWAKVPool(
|
self.token_to_kv_pool = SWAKVPool(
|
||||||
size=self.full_max_total_num_tokens,
|
size=self.full_max_total_num_tokens,
|
||||||
size_swa=self.swa_max_total_num_tokens,
|
size_swa=self.swa_max_total_num_tokens,
|
||||||
@@ -1997,7 +1996,7 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
if self.page_size == 1:
|
if self.page_size == 1:
|
||||||
if self.is_hybrid:
|
if self.is_hybrid_swa:
|
||||||
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
||||||
self.full_max_total_num_tokens,
|
self.full_max_total_num_tokens,
|
||||||
self.swa_max_total_num_tokens,
|
self.swa_max_total_num_tokens,
|
||||||
@@ -2015,7 +2014,7 @@ class ModelRunner:
|
|||||||
need_sort=need_sort,
|
need_sort=need_sort,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
assert not self.is_hybrid
|
assert not self.is_hybrid_swa
|
||||||
self.token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
|
self.token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
|
||||||
self.max_total_num_tokens,
|
self.max_total_num_tokens,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ class MockModelRunner:
|
|||||||
attention_arch = AttentionArch.MLA
|
attention_arch = AttentionArch.MLA
|
||||||
self.device = "cuda"
|
self.device = "cuda"
|
||||||
self.dtype = torch.float16
|
self.dtype = torch.float16
|
||||||
self.is_hybrid = False
|
self.is_hybrid_swa = False
|
||||||
context_len = 2048
|
context_len = 2048
|
||||||
self.model_config = type(
|
self.model_config = type(
|
||||||
"ModelConfig",
|
"ModelConfig",
|
||||||
|
|||||||
@@ -108,7 +108,7 @@ class MockModelRunner:
|
|||||||
self.config = {**DEFAULT_CONFIG, **(config or {})}
|
self.config = {**DEFAULT_CONFIG, **(config or {})}
|
||||||
self.dtype = self.config["dtype"]
|
self.dtype = self.config["dtype"]
|
||||||
self.kv_cache_dtype = self.config["kv_cache_dtype"]
|
self.kv_cache_dtype = self.config["kv_cache_dtype"]
|
||||||
self.is_hybrid = False
|
self.is_hybrid_swa = False
|
||||||
|
|
||||||
# Model configuration
|
# Model configuration
|
||||||
attention_arch = AttentionArch.MLA
|
attention_arch = AttentionArch.MLA
|
||||||
|
|||||||
Reference in New Issue
Block a user