[mem_cache] Route hybrid SWA full-side kv-row frees through free_segment (#37876)

Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
Co-authored-by: Sam Shleifer <sshleifer@gmail.com>
This commit is contained in:
Liangsheng Yin
2026-09-04 01:48:33 -07:00
committed by GitHub
co-authored by weireweire Sam Shleifer
parent dae126d510
commit 67248e04b4
11 changed files with 178 additions and 61 deletions
+18 -1
View File
@@ -192,6 +192,23 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
Starts sit on page boundaries, ends may fall mid-page, and the page Starts sit on page boundaries, ends may fall mid-page, and the page
ranges of consecutive segments do not overlap -- so in page units the ranges of consecutive segments do not overlap -- so in page units the
segments are aligned and disjoint, and every page is released once.""" segments are aligned and disjoint, and every page is released once."""
for free_index, start_pos in self._page_disjoint(segments):
self.free_segment(free_index, start_pos=start_pos)
def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int):
"""free_full() for a kv-row segment; same start-alignment contract as
free_segment(). Default: plain free_full()."""
assert start_pos % self.page_size == 0, (
f"segment start {start_pos} is not page-aligned"
)
self.free_full(free_index)
def free_full_segments(self, segments):
"""free_segments() for the full side alone; see free_full()."""
for free_index, start_pos in self._page_disjoint(segments):
self.free_full_segment(free_index, start_pos=start_pos)
def _page_disjoint(self, segments):
ps = self.page_size ps = self.page_size
prev_end = None prev_end = None
for free_index, start_pos in segments: for free_index, start_pos in segments:
@@ -202,4 +219,4 @@ class BaseTokenToKVPoolAllocator(abc.ABC):
f"segment at {start_pos} shares a page with the one ending at {prev_end}" f"segment at {start_pos} shares a page with the one ending at {prev_end}"
) )
prev_end = start_pos + n prev_end = start_pos + n
self.free_segment(free_index, start_pos=start_pos) yield free_index, start_pos
+31 -12
View File
@@ -103,7 +103,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.release_pages = None self.release_pages = None
self.free_group = None self.free_group = None
self.swa_free_group = [] self.swa_free_group = []
self.full_free_group = []
self._kvcache = kvcache self._kvcache = kvcache
self.clear() self.clear()
@@ -402,18 +401,34 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.full_to_swa_index_mapping[free_index] == 0, self.full_to_swa_index_mapping[free_index] == 0,
msg="caller wants free", msg="caller wants free",
) )
if self.free_group is None:
self.full_attn_allocator.free(free_index) self.full_attn_allocator.free(free_index)
else:
self.full_free_group.append(self._copy_for_free_group(free_index))
assert ( assert (
self.full_attn_allocator.available_size() <= self.full_attn_allocator.size self.full_attn_allocator.available_size() <= self.full_attn_allocator.size
) )
def free_segment(self, free_index: torch.Tensor, *, start_pos: int):
if free_index.numel() == 0:
return
# SWA first, as in free(): it reads the mapping that a later cache
# action in this group may re-point.
self.free_swa(free_index)
self.full_attn_allocator.free_segment(free_index, start_pos=start_pos)
def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int):
if free_index.numel() == 0:
return
expect(
_SWA_PEER_RELEASED,
self.full_to_swa_index_mapping[free_index] == 0,
msg="caller wants free_segment",
)
self.full_attn_allocator.free_segment(free_index, start_pos=start_pos)
def free_group_begin(self): def free_group_begin(self):
super().free_group_begin() super().free_group_begin()
self.swa_free_group = [] self.swa_free_group = []
self.full_free_group = [] # No full-side pile here: the full allocator's own group defers those.
self.full_attn_allocator.free_group_begin()
def free_group_end(self): def free_group_end(self):
super().free_group_end() super().free_group_end()
@@ -421,10 +436,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
swa_free_group = self.swa_free_group swa_free_group = self.swa_free_group
self.swa_free_group = [] self.swa_free_group = []
self._release_swa(torch.cat(swa_free_group)) self._release_swa(torch.cat(swa_free_group))
if self.full_free_group: self.full_attn_allocator.free_group_end()
full_free_group = self.full_free_group
self.full_free_group = []
self.full_attn_allocator.free(torch.cat(full_free_group))
assert ( assert (
self.full_attn_allocator.available_size() <= self.full_attn_allocator.size self.full_attn_allocator.available_size() <= self.full_attn_allocator.size
) )
@@ -468,7 +480,6 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
self.full_to_swa_index_mapping[:-1].fill_(0) self.full_to_swa_index_mapping[:-1].fill_(0)
self.free_group = None self.free_group = None
self.swa_free_group = [] self.swa_free_group = []
self.full_free_group = []
def get_cpu_copy(self, indices, mamba_indices=None): def get_cpu_copy(self, indices, mamba_indices=None):
return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices) return self._kvcache.get_cpu_copy(indices, mamba_indices=mamba_indices)
@@ -594,8 +605,16 @@ class PureSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
# release once the SWA side is gone. # release once the SWA side is gone.
return return
# Not inherited: the SWA parent's hooks drive swa_free_group and def free_segment(self, free_index: torch.Tensor, *, start_pos: int):
# full_free_group, which this pure-SWA variant does not have. # Single pool: the parent's split into an SWA and a full half would
# release the same slots twice.
self.free(free_index)
def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int):
return
# Not inherited: the SWA parent's hooks drive swa_free_group and the full
# allocator's group, which this pure-SWA variant does not have.
def free_group_begin(self): def free_group_begin(self):
BaseTokenToKVPoolAllocator.free_group_begin(self) BaseTokenToKVPoolAllocator.free_group_begin(self)
+6 -10
View File
@@ -127,7 +127,7 @@ def free_kv_row_segments(
) -> None: ) -> None:
"""Free ascending disjoint ``(kv_indices, start_pos)`` segments of one """Free ascending disjoint ``(kv_indices, start_pos)`` segments of one
request's kv row, split at the SWA eviction floor.""" request's kv row, split at the SWA eviction floor."""
swa_dead: list[torch.Tensor] = [] swa_dead: list[tuple[torch.Tensor, int]] = []
swa_alive: list[tuple[torch.Tensor, int]] = [] swa_alive: list[tuple[torch.Tensor, int]] = []
for kv_indices, start_pos in segments: for kv_indices, start_pos in segments:
num_indices = kv_indices.numel() num_indices = kv_indices.numel()
@@ -137,23 +137,19 @@ def free_kv_row_segments(
# the deliberately unmapped prefix of a PD decode SWA-tail prealloc. # the deliberately unmapped prefix of a PD decode SWA-tail prealloc.
num_dead = min(max(swa_evicted_seqlen - start_pos, 0), num_indices) num_dead = min(max(swa_evicted_seqlen - start_pos, 0), num_indices)
if num_dead > 0: if num_dead > 0:
swa_dead.append(kv_indices[:num_dead]) swa_dead.append((kv_indices[:num_dead], start_pos))
if num_dead < num_indices: if num_dead < num_indices:
swa_alive.append((kv_indices[num_dead:], start_pos + num_dead)) swa_alive.append((kv_indices[num_dead:], start_pos + num_dead))
if swa_dead and swa_alive: if swa_dead and swa_alive:
# A mid-page floor would send a page shared by the dead and alive # The two sides are separate calls, so neither one's page-disjointness
# sides back twice. # check sees a floor that splits a page between them.
assert swa_evicted_seqlen % allocator.page_size == 0, ( assert swa_evicted_seqlen % allocator.page_size == 0, (
f"SWA eviction floor {swa_evicted_seqlen} splits a page " f"SWA eviction floor {swa_evicted_seqlen} splits a page "
f"(page_size {allocator.page_size})" f"(page_size {allocator.page_size})"
) )
if len(swa_dead) == 1: if swa_dead:
allocator.free_full(swa_dead[0]) allocator.free_full_segments(swa_dead)
elif swa_dead:
# Two dead pieces can share a boundary page, and only free_full's own
# page dedup covers that -- free_segments trims the alive side alone.
allocator.free_full(torch.cat(swa_dead))
if swa_alive: if swa_alive:
allocator.free_segments(swa_alive) allocator.free_segments(swa_alive)
@@ -3227,6 +3227,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
self.free_group = None self.free_group = None
self.free_page_reps_group: Optional[List[torch.Tensor]] = None self.free_page_reps_group: Optional[List[torch.Tensor]] = None
self.full_free_group: List[torch.Tensor] = []
# Empty (not None) for the leak checker. # Empty (not None) for the leak checker.
self.free_pages = torch.empty(0, dtype=torch.int64, device=device) self.free_pages = torch.empty(0, dtype=torch.int64, device=device)
self.release_pages = torch.empty(0, dtype=torch.int64, device=device) self.release_pages = torch.empty(0, dtype=torch.int64, device=device)
@@ -3661,6 +3662,17 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
self.full_attn_allocator.free(free_index.detach().to(torch.int64)) self.full_attn_allocator.free(free_index.detach().to(torch.int64))
self.full_attn_allocator.clear_inverse_history() self.full_attn_allocator.clear_inverse_history()
def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None:
if free_index is None or free_index.numel() == 0:
return
if self.page_size == 1:
# token == page: free_full already frees by exact ids, no dedup.
self.free_full(free_index)
return
# The swa v2p is the mapping, so a tombstoned swa page drops out of the
# two-sided segment path by itself; full-only is the same call.
self.free_segment(free_index, start_pos=start_pos)
def set_full_to_swa_mapping( def set_full_to_swa_mapping(
self, full_indices: torch.Tensor, swa_indices: torch.Tensor self, full_indices: torch.Tensor, swa_indices: torch.Tensor
) -> None: ) -> None:
@@ -3676,13 +3688,20 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
# -- free-group -- # -- free-group --
# Not the SWA parent's hooks: those open the parent's paged full allocator
# as a free group, and this composite's sub-pools defer on their own.
def free_group_begin(self) -> None: def free_group_begin(self) -> None:
super().free_group_begin() BaseTokenToKVPoolAllocator.free_group_begin(self)
self.free_page_reps_group = [] self.free_page_reps_group = []
self.full_free_group = []
def free_group_end(self) -> None: def free_group_end(self) -> None:
pending, self.free_page_reps_group = self.free_page_reps_group, None pending, self.free_page_reps_group = self.free_page_reps_group, None
super().free_group_end() full_free_group, self.full_free_group = self.full_free_group, []
BaseTokenToKVPoolAllocator.free_group_end(self)
if full_free_group:
self.full_attn_allocator.free(torch.cat(full_free_group))
self.full_attn_allocator.clear_inverse_history()
if pending: if pending:
self._release_page_reps(pending) self._release_page_reps(pending)
@@ -3746,6 +3765,7 @@ class UnifiedSWATokenToKVPoolAllocator(SWATokenToKVPoolAllocator):
self.swa_attn_allocator.clear() self.swa_attn_allocator.clear()
self.free_group = None self.free_group = None
self.free_page_reps_group = None self.free_page_reps_group = None
self.full_free_group = []
# -- Lazy compaction hooks -- # -- Lazy compaction hooks --
@@ -503,10 +503,11 @@ class FullComponent(TreeComponent):
if isinstance(action, FreeComponentDeviceSlot): if isinstance(action, FreeComponentDeviceSlot):
alloc = self.cache.token_to_kv_pool_allocator alloc = self.cache.token_to_kv_pool_allocator
for indices in action.indices: for indices in action.indices:
# tree values are page-aligned copies of a kv row: page-exact segments
if self.cache.is_swa_enabled: if self.cache.is_swa_enabled:
alloc.full_attn_allocator.free(indices) alloc.full_attn_allocator.free_segment(indices, start_pos=0)
else: else:
alloc.free(indices) alloc.free_segment(indices, start_pos=0)
return return
raise AssertionError( raise AssertionError(
f"FullComponent: unhandled ComponentAction {type(action).__name__}" f"FullComponent: unhandled ComponentAction {type(action).__name__}"
@@ -1371,7 +1371,7 @@ class SWAComponent(TreeComponent):
swa_value = self._translate_full_to_swa(action.incoming_full) swa_value = self._translate_full_to_swa(action.incoming_full)
alloc.set_full_to_swa_mapping(action.kept_full, swa_value) alloc.set_full_to_swa_mapping(action.kept_full, swa_value)
alloc.clear_full_to_swa_mapping(action.incoming_full) alloc.clear_full_to_swa_mapping(action.incoming_full)
alloc.free_full(action.incoming_full) alloc.free_full_segment(action.incoming_full, start_pos=0)
self.tree_core.set_component_device_value( self.tree_core.set_component_device_value(
action.node_id, self.component_type, swa_value action.node_id, self.component_type, swa_value
) )
@@ -1097,7 +1097,7 @@ class UnifiedRadixCache(BasePrefixCache):
self.token_to_kv_pool_allocator.free_segment(indices, start_pos=0) self.token_to_kv_pool_allocator.free_segment(indices, start_pos=0)
elif isinstance(action, FreeDeviceKVFullOnly): elif isinstance(action, FreeDeviceKVFullOnly):
for indices in action.indices: for indices in action.indices:
self.token_to_kv_pool_allocator.free_full(indices) self.token_to_kv_pool_allocator.free_full_segment(indices, start_pos=0)
elif isinstance(action, BackupKV): elif isinstance(action, BackupKV):
if self.linker is not None: if self.linker is not None:
self.linker.offload_nodes(action.node_ids) self.linker.offload_nodes(action.node_ids)
@@ -141,7 +141,7 @@ class _StaticAllocRecorder:
self.clear_calls.append(full) self.clear_calls.append(full)
self.full_to_swa_index_mapping[full.to(torch.int64)] = 0 self.full_to_swa_index_mapping[full.to(torch.int64)] = 0
def free_full(self, indices): def free_full_segment(self, indices, *, start_pos):
self.freed_full.append(indices) self.freed_full.append(indices)
def free(self, indices): def free(self, indices):
@@ -307,8 +307,8 @@ class TestStaticPoolPathUnchanged(unittest.TestCase):
), ),
"incoming ids' mapping entries must be zeroed (static recipe)", "incoming ids' mapping entries must be zeroed (static recipe)",
) )
# Through free_full, not the inner allocator: the latter skips the # Through free_full_segment, not the inner allocator: the latter skips
# free-group defer. # the free-group defer.
self.assertEqual(len(static.freed_full), 1) self.assertEqual(len(static.freed_full), 1)
self.assertEqual(static.freed_via_inner, []) self.assertEqual(static.freed_via_inner, [])
((node_id, ct, _),) = probe.tree_core.set_calls ((node_id, ct, _),) = probe.tree_core.set_calls
@@ -2,6 +2,7 @@ import unittest
from array import array from array import array
from types import SimpleNamespace from types import SimpleNamespace
from unittest import mock from unittest import mock
from unittest.mock import patch
import torch import torch
@@ -830,13 +831,13 @@ class TestSWA(unittest.TestCase):
req2.prefix_indices = torch.tensor([21, 22, 23, 24, 25], device=tree.device) req2.prefix_indices = torch.tensor([21, 22, 23, 24, 25], device=tree.device)
freed_lens = [] freed_lens = []
original_free = allocator.free original_free_segment = allocator.free_segment
def wrapped_free(indices): def wrapped_free_segment(indices, *, start_pos):
freed_lens.append(int(indices.numel())) freed_lens.append(int(indices.numel()))
return original_free(indices) return original_free_segment(indices, start_pos=start_pos)
allocator.free = wrapped_free allocator.free_segment = wrapped_free_segment
tree.cache_finished_req( tree.cache_finished_req(
req2, is_insert=False, kv_len_to_handle=req2._kv_committed_len req2, is_insert=False, kv_len_to_handle=req2._kv_committed_len
) )
@@ -1070,21 +1071,45 @@ class TestFreeKvRow(CustomTestCase):
) )
self.assertEqual(self._sizes(), (self.full_baseline, self.swa_baseline)) self.assertEqual(self._sizes(), (self.full_baseline, self.swa_baseline))
def test_adjacent_below_floor_pieces_release_their_shared_page_once(self): def test_below_floor_pieces_go_back_through_the_full_side(self):
_, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4) _, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4)
indices = _swa_alloc(allocator, 8) indices = _swa_alloc(allocator, 8)
allocator.free_swa(indices) allocator.free_swa(indices)
after_alloc = allocator.full_available_size() after_alloc = allocator.full_available_size()
# Rows [0, 6) and [6, 8) both sit below the floor and share page 1. # Both rows [0, 4) and [4, 8) sit below the floor: full side only.
with patch.object(
allocator.full_attn_allocator,
"free",
side_effect=AssertionError("full side took the unique path"),
):
free_kv_row_segments( free_kv_row_segments(
allocator, allocator,
[(indices[:6], 0), (indices[6:], 6)], [(indices[:4], 0), (indices[4:], 4)],
swa_evicted_seqlen=8, swa_evicted_seqlen=8,
) )
self.assertEqual(allocator.full_available_size(), after_alloc + 8) self.assertEqual(allocator.full_available_size(), after_alloc + 8)
def test_grouped_full_side_frees_defer_and_skip_the_unique_path(self):
_, allocator, _ = _build_swa_tree(is_eagle=False, page_size=4)
indices = _swa_alloc(allocator, 12)
allocator.free_swa(indices[:8])
after_alloc = allocator.full_available_size()
with patch.object(
allocator.full_attn_allocator,
"free",
side_effect=AssertionError("full side took the unique path"),
):
allocator.free_group_begin()
# dead rows [0, 8) and the alive row [8, 12) from one request
free_kv_row_segments(allocator, [(indices, 0)], swa_evicted_seqlen=8)
self.assertEqual(allocator.full_available_size(), after_alloc)
allocator.free_group_end()
self.assertEqual(allocator.full_available_size(), after_alloc + 12)
def test_free_kv_row_reads_the_record_row_and_its_floor(self): def test_free_kv_row_reads_the_record_row_and_its_floor(self):
indices = _swa_alloc(self.allocator, 8) indices = _swa_alloc(self.allocator, 8)
cache = _RowCache(self.allocator, indices) cache = _RowCache(self.allocator, indices)
@@ -406,6 +406,26 @@ class TestEveryUnifiedAllocatorOverridesFreeSegment(unittest.TestCase):
self.assertIn("free_page_reps_group", inspect.getsource(cls)) self.assertIn("free_page_reps_group", inspect.getsource(cls))
class TestUnifiedSwaFullSideGroup(unittest.TestCase):
"""`free_full` inside a free group must defer and land at `free_group_end`;
the composite carries that pile itself, not through the SWA parent's hooks."""
def test_full_only_frees_defer_until_group_end(self):
from test_unified_swa_shared_virtual_ids import _build
alloc = _build(64, 32, 2, 4)
idx = alloc.alloc(8)
alloc.free_swa(idx)
before = alloc.full_available_size()
alloc.free_group_begin()
alloc.free_full(idx[:4])
alloc.free_full_segment(idx[4:], start_pos=4)
self.assertEqual(alloc.full_available_size(), before)
alloc.free_group_end()
self.assertEqual(alloc.full_available_size(), before + 8)
class TestFreeSwaWindowRatchetNoHostSync(unittest.TestCase): class TestFreeSwaWindowRatchetNoHostSync(unittest.TestCase):
"""The per-decode-step SWA window ratchet frees a CONTIGUOUS row slice """The per-decode-step SWA window ratchet frees a CONTIGUOUS row slice
with host-int, page-aligned bounds — the same shape `free_segment` was with host-int, page-aligned bounds — the same shape `free_segment` was
@@ -471,6 +491,25 @@ class TestFreeSwaWindowRatchetNoHostSync(unittest.TestCase):
alloc.free_swa(v[: 4 * self.PS], start_pos=0) alloc.free_swa(v[: 4 * self.PS], start_pos=0)
alloc.free_swa(v[4 * self.PS :], start_pos=4 * self.PS) alloc.free_swa(v[4 * self.PS :], start_pos=4 * self.PS)
def test_full_only_segment_free_never_syncs(self):
"""The request-finish dead half (swa already tombstoned) frees the full
side by page reps: no unique from free_full's token dedup."""
alloc = self._swa_composite(lazy=True)
v = alloc.alloc(8 * self.PS)
alloc.free_swa(v, start_pos=0)
before = alloc.full_available_size()
with (
mock.patch.object(
torch, "unique", side_effect=AssertionError("unique = host sync")
),
mock.patch.object(
torch.Tensor, "item", side_effect=AssertionError("item = host sync")
),
):
alloc.free_full_segment(v[: 4 * self.PS], start_pos=0)
alloc.free_full_segment(v[4 * self.PS :], start_pos=4 * self.PS)
self.assertEqual(alloc.full_available_size(), before + 8 * self.PS)
def test_unaligned_start_pos_is_rejected(self): def test_unaligned_start_pos_is_rejected(self):
"""A mid-page start must fail loudly, not release the head page whole.""" """A mid-page start must fail loudly, not release the head page whole."""
alloc = self._swa_composite(lazy=True) alloc = self._swa_composite(lazy=True)
@@ -516,6 +516,9 @@ def build_fixture(
full_attention_layer_ids=cfg.full_attention_layer_ids, full_attention_layer_ids=cfg.full_attention_layer_ids,
device=device, device=device,
) )
# Tree values reach the allocator as start_pos=0 segments, so only the
# debug cross-check against torch.unique catches a mis-aligned value.
with envs.SGLANG_DEBUG_MEMORY_POOL.override(True):
allocator = SWATokenToKVPoolAllocator( allocator = SWATokenToKVPoolAllocator(
size=cfg.kv_size, size=cfg.kv_size,
size_swa=cfg.kv_size, size_swa=cfg.kv_size,
@@ -4070,11 +4073,8 @@ class UnifiedRadixCacheSuite:
# for an un-charged gate to accept. # for an un-charged gate to accept.
extend_need = 2 * ps + 1 extend_need = 2 * ps + 1
max_new = 8 max_new = 8
self.assertIsNotNone( hold = cons_alloc.swa_available_size() - (window + extend_need - 1)
cons_alloc.swa_attn_allocator.alloc( self.assertIsNotNone(cons_alloc.swa_attn_allocator.alloc(hold // ps * ps))
cons_alloc.swa_available_size() - (window + extend_need - 1)
)
)
# The adder's SWA gate for this request (_swa_budget_for_req). # The adder's SWA gate for this request (_swa_budget_for_req).
surfaced_swa_hit = cons.staged_prefetch_swa_tokens(req_id) surfaced_swa_hit = cons.staged_prefetch_swa_tokens(req_id)
@@ -7559,8 +7559,8 @@ class TestUnifiedRadixCacheActionRouting(CustomTestCase):
_component_with_cache(ComponentType.FULL, cache).apply_component_action( _component_with_cache(ComponentType.FULL, cache).apply_component_action(
FreeComponentDeviceSlot([indices], component_type=ComponentType.FULL) FreeComponentDeviceSlot([indices], component_type=ComponentType.FULL)
) )
cache.token_to_kv_pool_allocator.full_attn_allocator.free.assert_called_once_with( cache.token_to_kv_pool_allocator.full_attn_allocator.free_segment.assert_called_once_with(
indices indices, start_pos=0
) )
def test_apply_component_action_device_kv_swa_uses_free_swa(self): def test_apply_component_action_device_kv_swa_uses_free_swa(self):
@@ -7681,7 +7681,7 @@ class TestUnifiedRadixCacheActionRouting(CustomTestCase):
alloc.set_full_to_swa_mapping.assert_called_once_with(kept_full, swa_value) alloc.set_full_to_swa_mapping.assert_called_once_with(kept_full, swa_value)
# the incoming full's stale mapping is cleared, then its slot freed (full-only) # the incoming full's stale mapping is cleared, then its slot freed (full-only)
alloc.clear_full_to_swa_mapping.assert_called_once_with(incoming_full) alloc.clear_full_to_swa_mapping.assert_called_once_with(incoming_full)
alloc.free_full.assert_called_once_with(incoming_full) alloc.free_full_segment.assert_called_once_with(incoming_full, start_pos=0)
# Never by indexing the tensor: the unified composite has no # Never by indexing the tensor: the unified composite has no
# `full_to_swa_index_mapping` to index into. # `full_to_swa_index_mapping` to index into.
alloc.full_to_swa_index_mapping.__setitem__.assert_not_called() alloc.full_to_swa_index_mapping.__setitem__.assert_not_called()