[mem_cache] Split duplicate insert frees at the SWA eviction floor (#37481)
This commit is contained in:
@@ -1071,9 +1071,17 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
|
||||
dup_start = max(0, state.params.prev_prefix_len - state.total_prefix_length)
|
||||
if dup_start < consumed_from:
|
||||
step_actions.append(
|
||||
FreeDeviceKV([value_slice[dup_start:consumed_from]])
|
||||
# The duplicate slice may straddle this request's own eviction
|
||||
# floor; below it only the full side is still ours to release.
|
||||
dup = value_slice[dup_start:consumed_from]
|
||||
abs_start = state.total_prefix_length + dup_start
|
||||
swa_already_freed = min(
|
||||
max(state.params.swa_evicted_seqlen - abs_start, 0), dup.numel()
|
||||
)
|
||||
if swa_already_freed > 0:
|
||||
step_actions.append(FreeDeviceKVFullOnly([dup[:swa_already_freed]]))
|
||||
if swa_already_freed < dup.numel():
|
||||
step_actions.append(FreeDeviceKV([dup[swa_already_freed:]]))
|
||||
|
||||
if self._inc_hit_count_and_check(node, state.params.chunked):
|
||||
step_actions.append(self._build_backup_kv_action(node))
|
||||
|
||||
@@ -851,11 +851,11 @@ fn insert_overlap_straddling_the_boundary_splits_and_recovers_the_tail() {
|
||||
node_id,
|
||||
source_value,
|
||||
},
|
||||
CacheAction::FreeDeviceKV(duplicates),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates),
|
||||
] = result.cache_actions.as_slice()
|
||||
else {
|
||||
panic!(
|
||||
"expected FreeDeviceKVFullOnly, SwaRebuild, FreeDeviceKV, got {:?}",
|
||||
"expected FreeDeviceKVFullOnly, SwaRebuild, FreeDeviceKVFullOnly, got {:?}",
|
||||
action_kinds(&result.cache_actions)
|
||||
);
|
||||
};
|
||||
@@ -898,10 +898,10 @@ fn insert_overlap_straddling_with_a_locked_full_defers_the_tail() {
|
||||
kept_full,
|
||||
incoming_full,
|
||||
},
|
||||
CacheAction::FreeDeviceKV(duplicates),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates),
|
||||
] = result.cache_actions.as_slice()
|
||||
else {
|
||||
panic!("expected RecoverSwaWithLockedFull then FreeDeviceKV");
|
||||
panic!("expected RecoverSwaWithLockedFull then FreeDeviceKVFullOnly");
|
||||
};
|
||||
assert_eq!(*node_id, tc.arena.node(node).id);
|
||||
assert!(kept_full.equal(&Tensor::from_slice(&[12i64, 13])));
|
||||
@@ -920,8 +920,8 @@ fn insert_overlap_entirely_outside_the_window_is_all_duplicate() {
|
||||
/* prev_prefix_len = */ 0,
|
||||
/* swa_evicted_seqlen = */ 3,
|
||||
));
|
||||
let [CacheAction::FreeDeviceKV(freed)] = result.cache_actions.as_slice() else {
|
||||
panic!("expected one FreeDeviceKV action");
|
||||
let [CacheAction::FreeDeviceKVFullOnly(freed)] = result.cache_actions.as_slice() else {
|
||||
panic!("expected one FreeDeviceKVFullOnly action");
|
||||
};
|
||||
assert!(freed[0].equal(&Tensor::from_slice(&[20i64, 21, 22])));
|
||||
}
|
||||
@@ -972,7 +972,7 @@ fn insert_overlap_boundary_at_the_node_start_recovers_the_whole_node() {
|
||||
.equal(&Tensor::from_slice(&[23i64, 24]))
|
||||
);
|
||||
let [
|
||||
CacheAction::FreeDeviceKV(duplicates),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates),
|
||||
CacheAction::FreeDeviceKVFullOnly(old_full),
|
||||
CacheAction::SwaRebuild {
|
||||
node_id,
|
||||
@@ -981,7 +981,7 @@ fn insert_overlap_boundary_at_the_node_start_recovers_the_whole_node() {
|
||||
] = result.cache_actions.as_slice()
|
||||
else {
|
||||
panic!(
|
||||
"expected FreeDeviceKV, FreeDeviceKVFullOnly, SwaRebuild, got {:?}",
|
||||
"expected FreeDeviceKVFullOnly, FreeDeviceKVFullOnly, SwaRebuild, got {:?}",
|
||||
action_kinds(&result.cache_actions)
|
||||
);
|
||||
};
|
||||
@@ -1027,17 +1027,17 @@ fn insert_overlap_straddling_a_second_level_node_recovers_the_tail() {
|
||||
.equal(&Tensor::from_slice(&[24i64]))
|
||||
);
|
||||
let [
|
||||
CacheAction::FreeDeviceKV(duplicates_head),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates_head),
|
||||
CacheAction::FreeDeviceKVFullOnly(old_tail),
|
||||
CacheAction::SwaRebuild {
|
||||
node_id,
|
||||
source_value,
|
||||
},
|
||||
CacheAction::FreeDeviceKV(duplicates_tail),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates_tail),
|
||||
] = result.cache_actions.as_slice()
|
||||
else {
|
||||
panic!(
|
||||
"expected FreeDeviceKV, FreeDeviceKVFullOnly, SwaRebuild, FreeDeviceKV, got {:?}",
|
||||
"expected FreeDeviceKVFullOnly, FreeDeviceKVFullOnly, SwaRebuild, FreeDeviceKVFullOnly, got {:?}",
|
||||
action_kinds(&result.cache_actions)
|
||||
);
|
||||
};
|
||||
@@ -1080,17 +1080,17 @@ fn insert_overlap_straddling_a_second_level_locked_node_defers_the_tail() {
|
||||
.equal(&Tensor::from_slice(&[14i64]))
|
||||
);
|
||||
let [
|
||||
CacheAction::FreeDeviceKV(duplicates_head),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates_head),
|
||||
CacheAction::RecoverSwaWithLockedFull {
|
||||
node_id,
|
||||
kept_full,
|
||||
incoming_full,
|
||||
},
|
||||
CacheAction::FreeDeviceKV(duplicates_tail),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates_tail),
|
||||
] = result.cache_actions.as_slice()
|
||||
else {
|
||||
panic!(
|
||||
"expected FreeDeviceKV, RecoverSwaWithLockedFull, FreeDeviceKV, got {:?}",
|
||||
"expected FreeDeviceKVFullOnly, RecoverSwaWithLockedFull, FreeDeviceKVFullOnly, got {:?}",
|
||||
action_kinds(&result.cache_actions)
|
||||
);
|
||||
};
|
||||
@@ -1321,7 +1321,7 @@ fn reinsert_boundary_at_the_node_start_rebuilds_the_whole_node() {
|
||||
));
|
||||
assert_eq!(tc.arena.node(b).key, vec![4, 5]);
|
||||
let [
|
||||
CacheAction::FreeDeviceKV(duplicates),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates),
|
||||
CacheAction::SwaRebuild {
|
||||
node_id,
|
||||
source_value,
|
||||
@@ -1329,7 +1329,7 @@ fn reinsert_boundary_at_the_node_start_rebuilds_the_whole_node() {
|
||||
] = result.cache_actions.as_slice()
|
||||
else {
|
||||
panic!(
|
||||
"expected FreeDeviceKV then SwaRebuild, got {:?}",
|
||||
"expected FreeDeviceKVFullOnly then SwaRebuild, got {:?}",
|
||||
action_kinds(&result.cache_actions)
|
||||
);
|
||||
};
|
||||
@@ -3137,11 +3137,11 @@ fn insert_overlap_straddling_with_a_partial_prev_prefix_recovers_the_tail() {
|
||||
node_id,
|
||||
source_value,
|
||||
},
|
||||
CacheAction::FreeDeviceKV(duplicates),
|
||||
CacheAction::FreeDeviceKVFullOnly(duplicates),
|
||||
] = result.cache_actions.as_slice()
|
||||
else {
|
||||
panic!(
|
||||
"expected FreeDeviceKVFullOnly, SwaRebuild, FreeDeviceKV, got {:?}",
|
||||
"expected FreeDeviceKVFullOnly, SwaRebuild, FreeDeviceKVFullOnly, got {:?}",
|
||||
action_kinds(&result.cache_actions)
|
||||
);
|
||||
};
|
||||
|
||||
@@ -1539,13 +1539,31 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
||||
|
||||
let dup_start = state.prev_prefix_len.saturating_sub(cursor);
|
||||
if dup_start < consumed_from {
|
||||
state
|
||||
.pending_actions
|
||||
.push(CacheAction::FreeDeviceKV(vec![value_slice.narrow(
|
||||
0,
|
||||
dup_start as i64,
|
||||
(consumed_from - dup_start) as i64,
|
||||
)]));
|
||||
// The duplicate slice may straddle this request's own eviction
|
||||
// floor; below it only the full side is still ours to release.
|
||||
let dup_len = consumed_from - dup_start;
|
||||
let swa_already_freed = state
|
||||
.swa_evicted_seqlen
|
||||
.saturating_sub(cursor + dup_start)
|
||||
.min(dup_len);
|
||||
if swa_already_freed > 0 {
|
||||
state
|
||||
.pending_actions
|
||||
.push(CacheAction::FreeDeviceKVFullOnly(vec![value_slice.narrow(
|
||||
0,
|
||||
dup_start as i64,
|
||||
swa_already_freed as i64,
|
||||
)]));
|
||||
}
|
||||
if swa_already_freed < dup_len {
|
||||
state.pending_actions.push(CacheAction::FreeDeviceKV(vec![
|
||||
value_slice.narrow(
|
||||
0,
|
||||
(dup_start + swa_already_freed) as i64,
|
||||
(dup_len - swa_already_freed) as i64,
|
||||
),
|
||||
]));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1292,7 +1292,9 @@ def test_swa_straddling_insert_crosses_the_boundary_actions():
|
||||
assert free_tail.indices[0].tolist() == [12, 13]
|
||||
assert isinstance(rebuild, SWARebuild)
|
||||
assert rebuild.source_value.tolist() == [22, 23]
|
||||
assert isinstance(free_duplicates, FreeDeviceKV)
|
||||
# Below the floor the duplicate's SWA peers are gone: full side only.
|
||||
assert isinstance(free_duplicates, FreeDeviceKVFullOnly)
|
||||
assert free_duplicates.indices[0].tolist() == [20, 21]
|
||||
|
||||
|
||||
def test_every_pool_name_crosses_the_prefetch_commit_boundary():
|
||||
|
||||
@@ -58,6 +58,7 @@ from sglang.srt.mem_cache.unified_cache.cache_action import (
|
||||
FreeComponentDeviceSlot,
|
||||
FreeComponentHostSlot,
|
||||
FreeDeviceKV,
|
||||
FreeDeviceKVFullOnly,
|
||||
RebuildFullToSWAMapping,
|
||||
RecoverSWAWithLockedFull,
|
||||
ReplaceWriteThroughOnNodeSplit,
|
||||
@@ -8096,6 +8097,30 @@ class TestResumableInsertWalkSWA(_InsertWalkSuite):
|
||||
self.assertIsNotNone(_device_value(cache, window_node, ComponentType.FULL))
|
||||
cache.sanity_check()
|
||||
|
||||
def test_dup_slice_below_eviction_floor_frees_full_only(self):
|
||||
"""A re-insert whose duplicate slice starts below the request's eviction
|
||||
floor gives back only the full side there: those SWA peers are gone."""
|
||||
sw = self.cfg.sliding_window_size
|
||||
cache, allocator, _ = build_fixture(self.cfg)
|
||||
seq = list(range(1, 2 * sw + 1))
|
||||
key = RadixKey(array("q", seq))
|
||||
cache.insert(InsertParams(key=key, value=self._alloc(allocator, len(seq))))
|
||||
|
||||
value = self._alloc(allocator, len(seq))
|
||||
with mock.patch.object(
|
||||
cache, "_apply_cache_action", wraps=cache._apply_cache_action
|
||||
) as spy:
|
||||
cache.insert(InsertParams(key=key, value=value, swa_evicted_seqlen=sw))
|
||||
actions = [c.args[0] for c in spy.call_args_list]
|
||||
|
||||
full_only = [
|
||||
i for a in actions if isinstance(a, FreeDeviceKVFullOnly) for i in a.indices
|
||||
]
|
||||
both = [i for a in actions if isinstance(a, FreeDeviceKV) for i in a.indices]
|
||||
torch.testing.assert_close(torch.cat(full_only), value[:sw])
|
||||
torch.testing.assert_close(torch.cat(both), value[sw:])
|
||||
cache.sanity_check()
|
||||
|
||||
def test_dec_swa_lock_only_early_release_keeps_full_lock(self):
|
||||
"""The scheduler's early SWA release (decode past the window) drops
|
||||
only the SWA lock; the Full path lock stays held until dec_lock_ref."""
|
||||
|
||||
Reference in New Issue
Block a user