[UnifiedTree]: Fix Unified HiCache tombstone lock release replay (#24972)
This commit is contained in:
@@ -4,7 +4,6 @@ import logging
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
|
||||
from sglang.srt.utils import get_bool_env_var
|
||||
|
||||
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
|
||||
@@ -701,16 +700,18 @@ class PrefillAdder:
|
||||
|
||||
@contextmanager
|
||||
def _lock_node(self, last_node: TreeNode):
|
||||
dec_lock_params = None
|
||||
try:
|
||||
result = self.tree_cache.inc_lock_ref(last_node)
|
||||
if self.tree_cache.supports_swa() and self.tree_cache.is_tree_cache():
|
||||
swa_uuid_for_lock = result.swa_uuid_for_lock
|
||||
if self.tree_cache.is_tree_cache():
|
||||
# init_load_back may revive SWA/Mamba tombstones while this
|
||||
# temporary admission lock is held. Release must mirror the
|
||||
# exact nodes skipped at acquire time.
|
||||
dec_lock_params = result.to_dec_params()
|
||||
yield None
|
||||
finally:
|
||||
if self.tree_cache.supports_swa() and self.tree_cache.is_tree_cache():
|
||||
self.tree_cache.dec_lock_ref(
|
||||
last_node, DecLockRefParams(swa_uuid_for_lock=swa_uuid_for_lock)
|
||||
)
|
||||
if dec_lock_params is not None:
|
||||
self.tree_cache.dec_lock_ref(last_node, dec_lock_params)
|
||||
else:
|
||||
self.tree_cache.dec_lock_ref(last_node)
|
||||
|
||||
|
||||
@@ -22,6 +22,9 @@ from sglang.srt.observability.metrics_collector import RadixCacheMetricsCollecto
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||
from sglang.srt.mem_cache.unified_cache_components.tree_component import (
|
||||
ComponentType,
|
||||
)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
@@ -94,10 +97,22 @@ class IncLockRefResult:
|
||||
|
||||
delta: Optional[int] = None
|
||||
swa_uuid_for_lock: Optional[int] = None
|
||||
# Component nodes that were tombstones at acquire time. Replaying this set
|
||||
# at release prevents a short-lived lock from consuming a later load-back or
|
||||
# request lock after that tombstone becomes a valid device value.
|
||||
skip_lock_node_ids: dict[ComponentType, set[int]] = dataclasses.field(
|
||||
default_factory=dict
|
||||
)
|
||||
|
||||
def to_dec_params(self) -> "DecLockRefParams":
|
||||
"""Convert to the corresponding DecLockRefParams for dec_lock_ref."""
|
||||
return DecLockRefParams(swa_uuid_for_lock=self.swa_uuid_for_lock)
|
||||
return DecLockRefParams(
|
||||
swa_uuid_for_lock=self.swa_uuid_for_lock,
|
||||
skip_lock_node_ids={
|
||||
component_type: set(node_ids)
|
||||
for component_type, node_ids in self.skip_lock_node_ids.items()
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
@@ -105,6 +120,9 @@ class DecLockRefParams:
|
||||
"""Parameters for dec_lock_ref operation."""
|
||||
|
||||
swa_uuid_for_lock: Optional[int] = None
|
||||
skip_lock_node_ids: dict[ComponentType, set[int]] = dataclasses.field(
|
||||
default_factory=dict
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
|
||||
@@ -166,9 +166,7 @@ class FullComponent(TreeComponent):
|
||||
cd.lock_ref += 1
|
||||
self.cache.evictable_device_leaves.discard(cur)
|
||||
cur = cur.parent
|
||||
result = IncLockRefResult(
|
||||
delta=delta, swa_uuid_for_lock=result.swa_uuid_for_lock
|
||||
)
|
||||
result.delta = delta
|
||||
return result
|
||||
|
||||
def release_component_lock(
|
||||
|
||||
@@ -217,12 +217,16 @@ class MambaComponent(TreeComponent):
|
||||
ct = self.component_type
|
||||
cd = node.component_data[ct]
|
||||
value = cd.value
|
||||
if value is not None:
|
||||
if cd.lock_ref == 0:
|
||||
vlen = len(value)
|
||||
self.cache.component_evictable_size_[ct] -= vlen
|
||||
self.cache.component_protected_size_[ct] += vlen
|
||||
cd.lock_ref += 1
|
||||
# A node in skip_lock_node_ids was a tombstone when this lock was acquired.
|
||||
if value is None:
|
||||
result.skip_lock_node_ids.setdefault(ct, set()).add(node.id)
|
||||
return result
|
||||
|
||||
if cd.lock_ref == 0:
|
||||
vlen = len(value)
|
||||
self.cache.component_evictable_size_[ct] -= vlen
|
||||
self.cache.component_protected_size_[ct] += vlen
|
||||
cd.lock_ref += 1
|
||||
return result
|
||||
|
||||
def release_component_lock(
|
||||
@@ -230,6 +234,10 @@ class MambaComponent(TreeComponent):
|
||||
) -> None:
|
||||
ct = self.component_type
|
||||
cd = node.component_data[ct]
|
||||
skip_lock_node_ids = params.skip_lock_node_ids.get(ct, ()) if params else ()
|
||||
if node.id in skip_lock_node_ids:
|
||||
return
|
||||
|
||||
value = cd.value
|
||||
if value is not None and cd.lock_ref > 0:
|
||||
if cd.lock_ref == 1:
|
||||
|
||||
@@ -362,6 +362,7 @@ class SWAComponent(TreeComponent):
|
||||
while cur != root and swa_lock_size < sliding_window_size:
|
||||
comp = cur.component_data[ct]
|
||||
if comp.value is None:
|
||||
result.skip_lock_node_ids.setdefault(ct, set()).add(cur.id)
|
||||
cur = cur.parent
|
||||
continue
|
||||
if comp.lock_ref == 0:
|
||||
@@ -385,14 +386,16 @@ class SWAComponent(TreeComponent):
|
||||
ct = self.component_type
|
||||
root = self.cache.root_node
|
||||
swa_uuid_for_lock = params.swa_uuid_for_lock if params else None
|
||||
skip_lock_node_ids = params.skip_lock_node_ids.get(ct, ()) if params else ()
|
||||
dec_swa = True
|
||||
|
||||
# lock_ref == 0 means acquire_component_lock skipped this node
|
||||
# (tombstone at acquire time) or load_back revived a tombstone between
|
||||
# acquire and release. Either way, there is nothing for us to undo here.
|
||||
# A node in skip_lock_node_ids was a tombstone when this lock was acquired.
|
||||
cur = node
|
||||
while cur != root and dec_swa:
|
||||
comp = cur.component_data[ct]
|
||||
if cur.id in skip_lock_node_ids:
|
||||
cur = cur.parent
|
||||
continue
|
||||
if comp.lock_ref == 0:
|
||||
cur = cur.parent
|
||||
continue
|
||||
|
||||
Reference in New Issue
Block a user