perf(unified-memory): batch lazy-compaction mapping lookup (#34066)
This commit is contained in:
@@ -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)."""
|
||||
|
||||
Reference in New Issue
Block a user