perf(unified-memory): batch lazy-compaction mapping lookup (#34066)

This commit is contained in:
SuperSong
2026-08-24 02:14:45 -07:00
committed by GitHub
parent 51b27f747a
commit d251fa2453
2 changed files with 55 additions and 27 deletions
@@ -1427,16 +1427,17 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
def _flush(self, *, urgent: bool) -> int:
"""One batched compaction pass; returns the number of survivor moves.
Pipeline (one D2H total, at step 3):
Pipeline (one free-list D2H plus one mapping D2H per committed move batch):
1. `_drain_pending_reuse` — return read-settled prior srcs.
2. sort the free list (or skip via env knob; either way ascending after).
3. `.tolist()` snapshot → `all_cpu` *(the one sync)*.
3. `.tolist()` snapshot → `all_cpu`.
4-5. `_absorb_boundary_holes` — retreat past boundary-contiguous holes;
`holes_cpu` = interior holes. After this `_free_phys_pages==holes_cpu`.
6. (urgent) `_settle_inflight_forward` — wait once so the walk is race-free.
7. survivor walk — TWO-POINTER: move topmost live slot into the next hole,
STOPPING when the pointers cross (band packed); batch into one
`move_kv_cache` + one v2p/p2v scatter at `_commit_move_batch`.
`move_kv_cache` + one v2p/p2v scatter at `_commit_move_batch`, which
gathers and validates all survivor virtual ids in one batch.
8-9. exit: urgent → FULL-PACK reclaim (retreat past ALL holes, empty list);
non-urgent → slice consumed dsts, merge freed srcs back.
@@ -1447,8 +1448,8 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
`_commit_move_batch` routes such srcs to `_pending_reuse`; urgent's
settle makes them immediately reusable.
`_topmost_survivor` excludes all p2v=-1 pages, so a `v_moved < 0` in the
loop is a corrupt-state bug and raises.
`_topmost_survivor` excludes all p2v=-1 pages, so a negative virtual id in
the batched mapping lookup is a corrupt-state bug and raises.
"""
if not self.lazy_compaction:
return 0
@@ -1463,7 +1464,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
if not _SORT_FREE_LIST_AFTER_MERGE and self._free_phys_pages.numel() > 1:
self._free_phys_pages, _ = torch.sort(self._free_phys_pages)
all_cpu = self._free_phys_pages.tolist() # the ONE D2H sync per flush
all_cpu = self._free_phys_pages.tolist() # one batched D2H sync
# `holes_cpu` = interior holes; `_free_phys_pages == holes_cpu` after.
new_wm, holes_cpu = self._absorb_boundary_holes(all_cpu)
@@ -1487,7 +1488,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
srcs: List[int] = []
dsts: List[int] = []
v_moveds: List[int] = []
# Flush-scoped accumulator for event-FIRED srcs. `_commit_move_batch`
# appends here instead of catting onto `_free_phys_pages`; the merge is
@@ -1532,12 +1532,11 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
# Commit accumulated moves, then wait the forward so the
# rest of the walk is race-free.
self._commit_move_batch(
srcs, dsts, v_moveds, latest_event, released_fired
srcs, dsts, latest_event, released_fired
)
n_moves += len(srcs)
srcs.clear()
dsts.clear()
v_moveds.clear()
inflight = self._inflight_forward
if inflight is not None:
torch.cuda.current_stream().wait_event(inflight[0])
@@ -1568,22 +1567,8 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
dst_cursor -= 1
n_dst_consumed += 1
v_moved = int(self.physical_to_virtual[src].item())
if v_moved < 0:
# `_topmost_survivor` excludes all p2v=-1 pages — corrupt state.
raise AssertionError(
f"MultiEndedAllocator({self.sub_pool_name!r})."
f"_flush: topmost survivor p={src} has p2v=-1; "
"this should be impossible (`_topmost_survivor` "
"excludes `holes_cpu` and `_pending_reuse_pages_cpu`)."
f" State: {self.allocator_state_str()}, "
f"#holes={len(holes_cpu)}, "
f"#pending_reuse={len(self._pending_reuse_pages_cpu)}"
)
srcs.append(src)
dsts.append(dst)
v_moveds.append(v_moved)
# Advance cursor strictly past the picked src.
if self.grow_direction == "up":
@@ -1594,7 +1579,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
if move_cap is not None and len(srcs) >= move_cap:
break
self._commit_move_batch(srcs, dsts, v_moveds, latest_event, released_fired)
self._commit_move_batch(srcs, dsts, latest_event, released_fired)
n_moves += len(srcs)
if single_pass_absorb:
@@ -1635,12 +1620,12 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
self,
srcs: List[int],
dsts: List[int],
v_moveds: List[int],
latest_event: Optional[torch.cuda.Event],
released_fired: List[torch.Tensor],
) -> None:
"""Issue ONE `move_kv_cache` + ONE bulk v2p/p2v remap for the accumulated
`(src, dst, v_moved)` triples. Fired srcs accumulate in `released_fired`
`(src, dst)` pairs. Survivor virtual ids are gathered from p2v in one
batch. Fired srcs accumulate in `released_fired`
(merged by `_flush` AFTER its dst-slice, keeping the free list == holes_cpu);
event-pending srcs route to `_pending_reuse` (read-race gating).
"""
@@ -1649,7 +1634,11 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
with record_function("MultiEndedAlloc._commit_move_batch"):
src_pages_t = torch.tensor(srcs, dtype=torch.int64, device=self.device)
dst_pages_t = torch.tensor(dsts, dtype=torch.int64, device=self.device)
v_moveds_t = torch.tensor(v_moveds, dtype=torch.int64, device=self.device)
v_moveds_t = self.physical_to_virtual[src_pages_t]
torch._assert_async(
(v_moveds_t >= 0).all(),
"invalid p2v mapping in MultiEndedAllocator._flush",
)
# Expand to token granularity (the move kernel is token-granular).
if self.page_size == 1:
src_t, dst_t = src_pages_t, dst_pages_t
@@ -79,6 +79,24 @@ class _FakeKVCache:
self.buf[dst_loc] = self.buf[src_loc].clone()
class _RejectScalarIndexTensor:
"""Tensor proxy that rejects one-row-at-a-time mapping lookups."""
def __init__(self, tensor: torch.Tensor):
self.tensor = tensor
def __getattr__(self, name):
return getattr(self.tensor, name)
def __getitem__(self, index):
if isinstance(index, int):
raise AssertionError("physical_to_virtual was read one row at a time")
return self.tensor[index]
def __setitem__(self, index, value):
self.tensor[index] = value
class TestUnifiedKVPoolViews(unittest.TestCase):
def test_min_slot_index_and_disjoint_bytes(self):
full = _make_mha_spec("full", "up", layer_num=4)
@@ -2049,6 +2067,27 @@ class TestLazyCompaction(unittest.TestCase):
continue # freed
self.assertEqual(int(fa.physical_to_virtual[p].item()), v)
def test_lazy_flush_gathers_survivor_mappings_as_one_batch(self):
"""Compaction must not synchronize once per relocated survivor."""
_pool, fa, kv = self._make_full(lazy=True)
values = fa.alloc(12)
self._stamp_kv(kv, fa, values)
fa.free(values[1:5].clone())
physical_to_virtual = fa.physical_to_virtual
fa.physical_to_virtual = _RejectScalarIndexTensor(physical_to_virtual)
try:
self.assertEqual(fa._flush(urgent=True), 4)
finally:
fa.physical_to_virtual = physical_to_virtual
for virtual in values.tolist():
physical = int(fa.virtual_to_physical[virtual].item())
if physical == -1:
continue
self.assertEqual(int(fa.physical_to_virtual[physical].item()), virtual)
self.assertEqual(int(kv.buf[physical].item()), virtual)
def _replay_sequence(self, ops, lazy: bool):
"""Run a given alloc/free op trace under eager OR lazy mode and
return the final (live virtual set, alloc-time KV stamps)."""