[UnifiedTree]: Fix Unified HiCache tombstone lock release replay (#24972)

This commit is contained in:
Zhangheng
2026-05-12 13:16:06 +08:00
committed by GitHub
parent 4ad63ad02f
commit 91907b7b93
6 changed files with 141 additions and 20 deletions
@@ -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