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:
|
def _flush(self, *, urgent: bool) -> int:
|
||||||
"""One batched compaction pass; returns the number of survivor moves.
|
"""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.
|
1. `_drain_pending_reuse` — return read-settled prior srcs.
|
||||||
2. sort the free list (or skip via env knob; either way ascending after).
|
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;
|
4-5. `_absorb_boundary_holes` — retreat past boundary-contiguous holes;
|
||||||
`holes_cpu` = interior holes. After this `_free_phys_pages==holes_cpu`.
|
`holes_cpu` = interior holes. After this `_free_phys_pages==holes_cpu`.
|
||||||
6. (urgent) `_settle_inflight_forward` — wait once so the walk is race-free.
|
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,
|
7. survivor walk — TWO-POINTER: move topmost live slot into the next hole,
|
||||||
STOPPING when the pointers cross (band packed); batch into one
|
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);
|
8-9. exit: urgent → FULL-PACK reclaim (retreat past ALL holes, empty list);
|
||||||
non-urgent → slice consumed dsts, merge freed srcs back.
|
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
|
`_commit_move_batch` routes such srcs to `_pending_reuse`; urgent's
|
||||||
settle makes them immediately reusable.
|
settle makes them immediately reusable.
|
||||||
|
|
||||||
`_topmost_survivor` excludes all p2v=-1 pages, so a `v_moved < 0` in the
|
`_topmost_survivor` excludes all p2v=-1 pages, so a negative virtual id in
|
||||||
loop is a corrupt-state bug and raises.
|
the batched mapping lookup is a corrupt-state bug and raises.
|
||||||
"""
|
"""
|
||||||
if not self.lazy_compaction:
|
if not self.lazy_compaction:
|
||||||
return 0
|
return 0
|
||||||
@@ -1463,7 +1464,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
if not _SORT_FREE_LIST_AFTER_MERGE and self._free_phys_pages.numel() > 1:
|
if not _SORT_FREE_LIST_AFTER_MERGE and self._free_phys_pages.numel() > 1:
|
||||||
self._free_phys_pages, _ = torch.sort(self._free_phys_pages)
|
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.
|
# `holes_cpu` = interior holes; `_free_phys_pages == holes_cpu` after.
|
||||||
new_wm, holes_cpu = self._absorb_boundary_holes(all_cpu)
|
new_wm, holes_cpu = self._absorb_boundary_holes(all_cpu)
|
||||||
@@ -1487,7 +1488,6 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
|
|
||||||
srcs: List[int] = []
|
srcs: List[int] = []
|
||||||
dsts: List[int] = []
|
dsts: List[int] = []
|
||||||
v_moveds: List[int] = []
|
|
||||||
|
|
||||||
# Flush-scoped accumulator for event-FIRED srcs. `_commit_move_batch`
|
# Flush-scoped accumulator for event-FIRED srcs. `_commit_move_batch`
|
||||||
# appends here instead of catting onto `_free_phys_pages`; the merge is
|
# 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
|
# Commit accumulated moves, then wait the forward so the
|
||||||
# rest of the walk is race-free.
|
# rest of the walk is race-free.
|
||||||
self._commit_move_batch(
|
self._commit_move_batch(
|
||||||
srcs, dsts, v_moveds, latest_event, released_fired
|
srcs, dsts, latest_event, released_fired
|
||||||
)
|
)
|
||||||
n_moves += len(srcs)
|
n_moves += len(srcs)
|
||||||
srcs.clear()
|
srcs.clear()
|
||||||
dsts.clear()
|
dsts.clear()
|
||||||
v_moveds.clear()
|
|
||||||
inflight = self._inflight_forward
|
inflight = self._inflight_forward
|
||||||
if inflight is not None:
|
if inflight is not None:
|
||||||
torch.cuda.current_stream().wait_event(inflight[0])
|
torch.cuda.current_stream().wait_event(inflight[0])
|
||||||
@@ -1568,22 +1567,8 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
dst_cursor -= 1
|
dst_cursor -= 1
|
||||||
n_dst_consumed += 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)
|
srcs.append(src)
|
||||||
dsts.append(dst)
|
dsts.append(dst)
|
||||||
v_moveds.append(v_moved)
|
|
||||||
|
|
||||||
# Advance cursor strictly past the picked src.
|
# Advance cursor strictly past the picked src.
|
||||||
if self.grow_direction == "up":
|
if self.grow_direction == "up":
|
||||||
@@ -1594,7 +1579,7 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
if move_cap is not None and len(srcs) >= move_cap:
|
if move_cap is not None and len(srcs) >= move_cap:
|
||||||
break
|
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)
|
n_moves += len(srcs)
|
||||||
|
|
||||||
if single_pass_absorb:
|
if single_pass_absorb:
|
||||||
@@ -1635,12 +1620,12 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
self,
|
self,
|
||||||
srcs: List[int],
|
srcs: List[int],
|
||||||
dsts: List[int],
|
dsts: List[int],
|
||||||
v_moveds: List[int],
|
|
||||||
latest_event: Optional[torch.cuda.Event],
|
latest_event: Optional[torch.cuda.Event],
|
||||||
released_fired: List[torch.Tensor],
|
released_fired: List[torch.Tensor],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Issue ONE `move_kv_cache` + ONE bulk v2p/p2v remap for the accumulated
|
"""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);
|
(merged by `_flush` AFTER its dst-slice, keeping the free list == holes_cpu);
|
||||||
event-pending srcs route to `_pending_reuse` (read-race gating).
|
event-pending srcs route to `_pending_reuse` (read-race gating).
|
||||||
"""
|
"""
|
||||||
@@ -1649,7 +1634,11 @@ class MultiEndedAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
with record_function("MultiEndedAlloc._commit_move_batch"):
|
with record_function("MultiEndedAlloc._commit_move_batch"):
|
||||||
src_pages_t = torch.tensor(srcs, dtype=torch.int64, device=self.device)
|
src_pages_t = torch.tensor(srcs, dtype=torch.int64, device=self.device)
|
||||||
dst_pages_t = torch.tensor(dsts, 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).
|
# Expand to token granularity (the move kernel is token-granular).
|
||||||
if self.page_size == 1:
|
if self.page_size == 1:
|
||||||
src_t, dst_t = src_pages_t, dst_pages_t
|
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()
|
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):
|
class TestUnifiedKVPoolViews(unittest.TestCase):
|
||||||
def test_min_slot_index_and_disjoint_bytes(self):
|
def test_min_slot_index_and_disjoint_bytes(self):
|
||||||
full = _make_mha_spec("full", "up", layer_num=4)
|
full = _make_mha_spec("full", "up", layer_num=4)
|
||||||
@@ -2049,6 +2067,27 @@ class TestLazyCompaction(unittest.TestCase):
|
|||||||
continue # freed
|
continue # freed
|
||||||
self.assertEqual(int(fa.physical_to_virtual[p].item()), v)
|
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):
|
def _replay_sequence(self, ops, lazy: bool):
|
||||||
"""Run a given alloc/free op trace under eager OR lazy mode and
|
"""Run a given alloc/free op trace under eager OR lazy mode and
|
||||||
return the final (live virtual set, alloc-time KV stamps)."""
|
return the final (live virtual set, alloc-time KV stamps)."""
|
||||||
|
|||||||
Reference in New Issue
Block a user