Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions docs/environment_variables.rst
Original file line number Diff line number Diff line change
Expand Up @@ -244,3 +244,9 @@ Worker scheduling
- ``0``
- ``N > 0``: every N iterations log per-phase p50/p95/mean of the
worker main loop (speculate, await_gpu, submit_spec, ...).
* - ``MSTAR_KV_DEBUG_ASSERTS``

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nit: not a fan of this flag existing but it doesn't hurt ability, although does disperse testing across a wider than already-is range

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The tests don't need it beacuse they call assert_pages_conserved directly. I kept it for checking a live server. I'm currently stress testing #280 with shared prefixes, cancellations and eviction, and with the flag on, a page that leaks or gets freed while still in use raises at the op that caused it.

- ``0``
- ``1``: after every ``admit``, ``commit``, ``reset_request`` and
``remove_request``, check the KV page bookkeeping (free list, owner
counts, seals) against the streams holding the pages. Walks every
live stream; tests and debugging only.
117 changes: 112 additions & 5 deletions mstar/engine/resources/kv/manager.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import logging
import os
import threading
from concurrent.futures import Future, wait
from dataclasses import dataclass, field
Expand Down Expand Up @@ -39,18 +40,56 @@

logger = logging.getLogger(__name__)

# Off by default: `assert_pages_conserved` walks every stream of every live
# request after each admit, commit, reset and remove.
_DEBUG_ASSERTS = os.environ.get("MSTAR_KV_DEBUG_ASSERTS", "0") == "1"


@dataclass
class PageArena:
"""physical storage and free list management"""
"""physical storage, free list, and per-page ownership

A page goes back to the allocator when its last owner releases it, not its
first. A sealed page is never written again, so a second owner may read it;
freeing clears the seal. Every caller holds the manager's `_lock`, so the
counts and seals need none of their own.
"""
kv_cache: KVCache
allocator: PageAllocator
num_owners: list[int] = field(init=False, repr=False)
sealed: list[bool] = field(init=False, repr=False)

def __post_init__(self):
self.num_owners = [0] * self.allocator.max_num_pages
self.sealed = [False] * self.allocator.max_num_pages

def acquire(self, n: int) -> list[int] | None:
return self.allocator.try_allocate(n)
pages = self.allocator.try_allocate(n)
if pages is not None:
for page in pages:
self.num_owners[page] = 1
return pages

def retain(self, pages: list[int]) -> None:
for page in pages:
self.num_owners[page] += 1

def seal(self, pages: list[int]) -> None:
for page in pages:
self.sealed[page] = True

def any_sealed(self, pages: list[int]) -> bool:
return any(self.sealed[page] for page in pages)

def release(self, pages: list[int]) -> None:
return self.allocator.free(pages)
freed = []
for page in pages:
assert self.num_owners[page] > 0, f"page {page} released with no owner"
self.num_owners[page] -= 1
if self.num_owners[page] == 0:
self.sealed[page] = False
freed.append(page)
self.allocator.free(freed)

def copy_pages(self, src: list[int], dst: list[int]) -> None:
self.kv_cache.copy_pages(src, dst)
Expand Down Expand Up @@ -431,6 +470,8 @@ def admit(self, step: KVStep, ctx: StepContext) -> AdmitOutcome:
self._preplan_marked.append(
(segment.request_id, segment.label)
)
if _DEBUG_ASSERTS:
self.assert_pages_conserved()
# TODO: apply retention policy

return ADMIT_OK
Expand Down Expand Up @@ -665,6 +706,8 @@ def commit(self, step: KVStep, ctx: StepContext):
for (from_label, to_label) in step.post_forks:
for rid in ctx.padded_request_ids:
self._apply_fork(rid, from_label, to_label)
if _DEBUG_ASSERTS:
self.assert_pages_conserved()
# TODO: handle retention policy, free pages if not commit

# Eviction
Expand Down Expand Up @@ -906,9 +949,15 @@ def reset_request(self, rid: str, free: bool=False):
wait([stream.read_future])
with self._lock:
for stream in self._streams.get(rid, {}).values():
if free:
# a rewind would put the next write on the stream's first page,
# over a sealed one its other owners still read. drop the pages
# and let the next write allocate; `free` asks for the same
drop = free or self._arena.any_sealed(stream.page_indices)
if drop:
self._arena.release(stream.page_indices)
stream.reset(freed=free)
stream.reset(freed=drop)
if _DEBUG_ASSERTS:
self.assert_pages_conserved()

def remove_request(self, rid: str):
streams = self._streams.get(rid)
Expand All @@ -925,6 +974,64 @@ def remove_request(self, rid: str):
self._cpu_pool.remove_request(rid)
self._streams.pop(rid, None)
self._overrides.pop(rid, None)
if _DEBUG_ASSERTS:
self.assert_pages_conserved()

def assert_pages_conserved(self) -> None:
"""Check the owner counts against the streams holding the pages.

Each assertion names the rule it checks. Host pages are not covered:
`CPUPagePool` keeps no counts.
"""
with self._lock:
arena = self._arena
free = list(arena.allocator.free_pages.queue)
owned = [
page for page in range(self.config.max_num_pages)
if arena.num_owners[page] > 0
]
both = sorted(set(free) & set(owned))
assert not both, f"pages both free and owned: {both}"
assert len(free) + len(owned) == self.config.max_num_pages, (
f"{self.config.max_num_pages} pages in the pool, but "
f"{len(free)} free and {len(owned)} owned"
)

# the sink belongs to no request, so count it by hand. `frontier`
# is the page each stream is still writing into, plus any it holds
# past that
refs: dict[int, int] = {SINK_PAGE: 1}
frontier: set[int] = set()
for streams in self._streams.values():
for stream in streams.values():
for page in stream.page_indices:
refs[page] = refs.get(page, 0) + 1
full = stream.stored_len // self.config.page_size
frontier.update(stream.page_indices[full:])

counts = {page: arena.num_owners[page] for page in owned}
assert counts == refs, (
"owner counts disagree with the streams naming the pages: "
+ ", ".join(
f"page {page} owned {counts.get(page, 0)}, "
f"named {refs.get(page, 0)}"
for page in sorted(set(counts) | set(refs))
if counts.get(page, 0) != refs.get(page, 0)
)
)

unsealed = [
page for page in owned
if arena.num_owners[page] > 1 and not arena.sealed[page]
]
assert not unsealed, f"pages shared before they were sealed: {unsealed}"

crowded = sorted(
page for page in frontier if arena.num_owners[page] != 1
)
assert not crowded, (
f"pages still being written into, but not owned alone: {crowded}"
)

def post_warmup_validate(self):
"""Assert ``num_free_pages`` is identical across every TP rank
Expand Down
Loading
Loading