Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 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
1 change: 1 addition & 0 deletions docs/design_overview.md
Original file line number Diff line number Diff line change
Expand Up @@ -198,6 +198,7 @@ These statuses describe current KVCR knowledge, not a reservation or guarantee.
`deposit` copies from framework-owned memory into the KVCR's pool, while `deliver` places data into a framework-provided destination. `deliver` does not name a source; source selection remains with the KVCR and router. The KVCR does not allocate or free framework memory.

To serve from framework-owned memory, the KVCR acquires a pin asynchronously through `request_pin` and `poll_pin_results`, reusing covered keys and requesting only the remainder; the framework keeps it valid until the KVCR calls `release_pin`.
`release_pin` must be safely retryable: `False` or an exception leaves release pending; `True` means the framework accepts responsibility for completing release.

A deployment may choose to use only framework-owned memory. In that case, it uses the pinning mechanism together with `deliver` and does not use `deposit`, `fetch`, or `release`.

Expand Down
2 changes: 2 additions & 0 deletions src/kvcr/control_channels.py
Original file line number Diff line number Diff line change
Expand Up @@ -295,6 +295,8 @@ def recv(self) -> list[bytes]:
except zmq.Again:
return messages
except zmq.ZMQError:
# Warning suppression can be added if persistent ZMQ failures
# cause excessive receive logs.
logger.warning("KVCR control recv failed", exc_info=True)
return messages

Expand Down
5 changes: 4 additions & 1 deletion src/kvcr/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -773,7 +773,10 @@ def _discard_local_dram_fill(self, keys: Collection[BlockKey]) -> None:
self._local_dram.discard_fill(keys)

def _block_record(self, key: BlockKey) -> _BlockRecord:
return self._block_record_map.setdefault(key, _BlockRecord())
record = self._block_record_map.get(key)
if record is None:
record = self._block_record_map[key] = _BlockRecord()
return record

def _is_local_resident(self, key: BlockKey) -> bool:
"""Report local DRAM residency a new operation can still be served from.
Expand Down
28 changes: 13 additions & 15 deletions src/kvcr/local_dram.py
Original file line number Diff line number Diff line change
Expand Up @@ -432,14 +432,9 @@ def complete_fill(self, keys: Collection[BlockKey], *, success: bool) -> None:
for key in ordered_keys:
record = self._kvcr._block_record_map.get(key)
residency = record.local_dram if record is not None else None
if (
residency is None
or residency.state
not in (
_LocalDramState.FILLING,
_LocalDramState.DISCARDING,
)
or (success and residency.state is not _LocalDramState.FILLING)
if residency is None or residency.state not in (
_LocalDramState.FILLING,
_LocalDramState.DISCARDING,
):
raise RuntimeError(f"local DRAM fill state lost for {key!r}")
slots.append(tuple(residency.slots))
Expand Down Expand Up @@ -605,6 +600,7 @@ def _apply_fill_result(
source: CacheTier,
) -> None:
committed: list[BlockKey] = []
failed: list[BlockKey] = []
affected_residency_ops: dict[_OpId, _PendingResidencyOp] = {}
affected_deliver_ops: dict[_OpId, _PendingDeliverOp] = {}
deliver_keys: dict[_OpId, list[BlockKey]] = {}
Expand All @@ -621,10 +617,12 @@ def _apply_fill_result(
_LocalDramState.FILLING,
_LocalDramState.DISCARDING,
)
or (success and residency.state is not _LocalDramState.FILLING)
):
raise RuntimeError(f"local DRAM fill state lost for {key!r}")
if success:
# Main may discard a fill after progress queues its success.
# A terminal completion then frees the slot instead of committing it.
key_success = success and residency.state is _LocalDramState.FILLING
if key_success:
record.last_access = now
residency.state = _LocalDramState.READY
self._residency_observer(key, record)
Expand All @@ -637,11 +635,12 @@ def _apply_fill_result(
else:
record.local_dram = None
self._free(residency.slots)
failed.append(key)

for op_id in record.active_op_ids:
residency_op = self._pending_residency_ops.get(op_id)
if residency_op is not None and key in residency_op.keys:
if success and (
if key_success and (
residency_op.op_id[0] == "deposit"
or now < residency_op.deadline
):
Expand All @@ -665,7 +664,7 @@ def _apply_fill_result(

deliver_op = self._pending_deliver_ops.get(op_id)
if deliver_op is not None and key in deliver_op.keys:
if success:
if key_success:
deliver_keys.setdefault(op_id, []).append(key)
else:
deliver_op.results[key] = OpEntryResult(OpEntryStatus.FAILED)
Expand All @@ -677,9 +676,8 @@ def _apply_fill_result(
self._finish_residency_if_ready(residency_op)
for op_id, deliver_op in affected_deliver_ops.items():
self._start_deliveries(deliver_op, deliver_keys.get(op_id, ()))
if not success:
for key in ordered_keys:
self._kvcr._prune_block_record(key)
for key in failed:
self._kvcr._prune_block_record(key)
self._resume_capacity_waiters()

def reserve_fill(
Expand Down
3 changes: 3 additions & 0 deletions src/kvcr/policy_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -151,6 +151,9 @@ class _Entry:


class _EvictionQueue:
# TODO: Bound stale heap growth from DRAM/G3 claim/release cycles.
# Removal only invalidates _live; select() removes stale heap entries as
# it encounters them, so repeated cache use can grow _heap without eviction.
def __init__(self) -> None:
self._heap: list[tuple[float, int, BlockKey]] = []
self._live: dict[BlockKey, _Entry] = {}
Expand Down
32 changes: 21 additions & 11 deletions src/kvcr/progress.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,8 +19,9 @@

logger = logging.getLogger(__name__)
_IDLE_WAIT_SECONDS = 0.001
_CLOSE_TIMEOUT_SECONDS = 5.0
_OP_CLEANUP_TIMEOUT_SECONDS = 5.0
_JOIN_TIMEOUT_SECONDS = 10.0
_STARTUP_TIMEOUT_SECONDS = 30.0
_RELEASE_LOG_INTERVAL_SECONDS = 1.0
_STOP = object()
_OpId = tuple[str, Any]
Expand Down Expand Up @@ -108,6 +109,7 @@ def __init__(
)
self._failure: BaseException | None = None
self._stop_requested = False
self._startup_stage = "thread startup"

@property
def nixl_agent(self) -> Any:
Expand Down Expand Up @@ -256,8 +258,11 @@ def _release_transfer(self, transfer_id: int, state: _TransferState) -> bool:

def start(self) -> None:
self._thread.start()
if not self._ready.wait(timeout=_JOIN_TIMEOUT_SECONDS):
raise RuntimeError("KVCR progress thread did not start")
if not self._ready.wait(timeout=_STARTUP_TIMEOUT_SECONDS):
raise RuntimeError(
"KVCR progress initialization timed out after "
f"{_STARTUP_TIMEOUT_SECONDS:g}s (stage: {self._startup_stage})"
)
self.raise_if_failed()

def submit(self, item: object) -> None:
Expand Down Expand Up @@ -302,19 +307,25 @@ def close(self) -> None:

def _run(self) -> None:
try:
self._startup_stage = "NIXL agent initialization"
self._initialize_nixl()
# Let KVCR backends initialize NIXL resources before common
# memory registration.
self._startup_stage = "backend initialization"
self._initialize(self)
self._startup_stage = "memory registration"
self._register_memory_regions()
self._startup_stage = "agent metadata capture"
self._capture_agent_metadata()
self._startup_stage = "ready"
self._ready.set()
while not self._stop_requested:
if not self._run_one_iteration():
time.sleep(_IDLE_WAIT_SECONDS)
except BaseException as error:
self._failure = error
finally:
self._startup_stage = "cleanup"
try:
try:
self._close_progress_ops()
Expand All @@ -329,7 +340,7 @@ def _run(self) -> None:
self._ready.set()

def _close_progress_ops(self) -> None:
deadline = time.monotonic() + _CLOSE_TIMEOUT_SECONDS
deadline = time.monotonic() + _OP_CLEANUP_TIMEOUT_SECONDS
while self._in_flight_ops:
for op_id, op in list(self._in_flight_ops.items()):
if op.close(self):
Expand Down Expand Up @@ -397,15 +408,14 @@ def _initialize_nixl(self) -> None:
)

def _register_memory_regions(self) -> None:
if self._nixl_agent is None:
if self._nixl_agent is None or not self._memory_regions:
return
for address, size in self._memory_regions:
self._memory_registrations.append(
self._nixl_agent.register_memory(
[(address, size, 0, "")],
mem_type="DRAM",
)
self._memory_registrations.append(
self._nixl_agent.register_memory(
[(address, size, 0, "") for address, size in self._memory_regions],
mem_type="DRAM",
)
)

def _capture_agent_metadata(self) -> None:
get_agent_metadata = getattr(self._nixl_agent, "get_agent_metadata", None)
Expand Down
58 changes: 41 additions & 17 deletions src/kvcr/remote_fw_dram.py
Original file line number Diff line number Diff line change
Expand Up @@ -411,6 +411,7 @@ def __init__(
self._pending_pin_keys: dict[BlockKey, set[PinRequestId]] = {}
# Pins retained while their source operations execute.
self._fw_pins_by_op: dict[_OpId, set[PinHandle]] = {}
self._releasing_framework_pins: set[PinHandle] = set()

# Progress-thread state: G2 control and outbound events.
self._progress_outbound: list[object] = []
Expand Down Expand Up @@ -560,6 +561,7 @@ def fetch(
)

def poll_main(self, items: Collection[object]) -> None:
self._release_framework_pins(tuple(self._releasing_framework_pins))
for item in items:
if isinstance(item, _SourcePinOp):
self._start_source_pin(item)
Expand Down Expand Up @@ -603,6 +605,8 @@ def poll_main(self, items: Collection[object]) -> None:
if now >= op.deadline:
logger.warning("KVCR operation %r expired", op_id)
self._expire_source_pin(op_id, op)
elif not op.pending_pin_ids:
self._resume_source_pin(op_id, op)

def discard_hint(self, request_id: str) -> None:
self._request_hints.pop(request_id, None)
Expand Down Expand Up @@ -700,6 +704,8 @@ def poll_progress(
else:
raise TypeError(f"unsupported KVCR progress item: {type(item)!r}")

# Warning suppression can be added if persistent backend faults
# cause excessive polling logs.
observed_work |= self._process_control_messages(progress)
# A real notification outranks a refusal for the same operation.
events = {**self._refused_writes, **self._poll_notifications(progress)}
Expand Down Expand Up @@ -973,7 +979,7 @@ def _submit_prepared_source_write(
and record.fw_mem is not None
}
sources = {} if force_failure else {**framework_sources, **local_sources}
completed_keys: tuple[BlockKey, ...] = ()
completed_count = 0
for index, key in enumerate(source_pin.ordered_keys):
source = sources.get(key)
destination = source_pin.dst_descriptors[index]
Expand All @@ -988,7 +994,8 @@ def _submit_prepared_source_write(
key,
)
break
completed_keys = source_pin.ordered_keys[: index + 1]
completed_count += 1
completed_keys = source_pin.ordered_keys[:completed_count]

kvcr._release_local_dram_sources(
source_pin.op_id, local_sources.keys() - set(completed_keys)
Expand Down Expand Up @@ -1083,15 +1090,15 @@ def _process_pending_pin_results(self) -> None:
and request in op.pending_pin_ids
]
if not ops:
self._discard_pin_result(result)
self._discard_pin_result(result, wait.keys if wait is not None else ())
continue

now = kvcr._clock()
expired_ops = [(op_id, op) for op_id, op in ops if now >= op.deadline]
active_ops = [(op_id, op) for op_id, op in ops if now < op.deadline]
if not active_ops:
self._record_pending_pin_wait(wait, "timeout")
self._discard_pin_result(result)
self._discard_pin_result(result, wait.keys)
for op_id, op in expired_ops:
op.pending_pin_ids.discard(request)
self._expire_source_pin(op_id, op)
Expand Down Expand Up @@ -1165,6 +1172,11 @@ def _resume_source_pin(self, op_id: _OpId, op: _SourcePinOp) -> None:
key for key in op.ordered_keys if key not in local_sources
)
if unresolved_keys and not op.framework_acquire_attempted:
if any(
not kvcr._framework_pin_keys[pin].isdisjoint(unresolved_keys)
for pin in self._releasing_framework_pins
):
return
op.framework_acquire_attempted = True
framework_sources = self._acquire_framework_sources(unresolved_keys)
if isinstance(framework_sources, _PendingFrameworkSources):
Expand Down Expand Up @@ -1270,15 +1282,21 @@ def _record_pending_pin_wait(
if wait is not None:
self._kvcr._record_duration("framework_pin_wait", wait.started_at, result)

def _discard_pin_result(self, result: PinResult) -> None:
def _discard_pin_result(
self, result: PinResult, keys: Collection[BlockKey] = ()
) -> None:
if result is None:
return
try:
pin_handle = result[0]
except (IndexError, TypeError):
return
if isinstance(pin_handle, str):
self._try_release_pin(pin_handle)
pin_keys = self._kvcr._framework_pin_keys.setdefault(pin_handle, set())
pin_keys.update(keys)
if len(result) > 1 and isinstance(result[1], Mapping):
pin_keys.update(result[1])
self._release_framework_pins((pin_handle,))

# Framework pin ownership.

Expand Down Expand Up @@ -1310,7 +1328,6 @@ def _install_framework_pin(
keys: Collection[BlockKey],
pin_result: tuple[PinHandle, Mapping[BlockKey, list[MemDescriptor] | None]],
) -> PinHandle | None:
pin_handle: PinHandle | None = None
try:
pin_handle, descriptors = pin_result
if not isinstance(pin_handle, str) or not isinstance(descriptors, Mapping):
Expand Down Expand Up @@ -1340,8 +1357,7 @@ def _install_framework_pin(
pin_keys.add(key)
return pin_handle
except Exception:
if pin_handle is not None:
self._try_release_pin(pin_handle)
self._discard_pin_result(pin_result, keys)
return None

def _acquire_framework_sources(
Expand Down Expand Up @@ -1404,9 +1420,10 @@ def _release_framework_pins(self, framework_pins: Collection[PinHandle]) -> None
pin_keys = kvcr._framework_pin_keys.get(pin_handle)
if pin_keys is None:
continue
if not self._try_release_pin(pin_handle):
continue
kvcr._framework_pin_keys.pop(pin_handle, None)
# Once release starts, its descriptors are no longer safe to reuse.
# Keep the keys until release is accepted, so overlapping pins wait.
first_attempt = pin_handle not in self._releasing_framework_pins
self._releasing_framework_pins.add(pin_handle)
for key in pin_keys:
record = kvcr._block_record_map.get(key)
if (
Expand All @@ -1416,17 +1433,24 @@ def _release_framework_pins(self, framework_pins: Collection[PinHandle]) -> None
):
record.fw_mem = None
kvcr._prune_block_record(key)
if self._try_release_pin(pin_handle, warn=first_attempt):
kvcr._framework_pin_keys.pop(pin_handle, None)
self._releasing_framework_pins.discard(pin_handle)

def _try_release_pin(self, pin_handle: PinHandle) -> bool:
def _try_release_pin(self, pin_handle: PinHandle, *, warn: bool) -> bool:
# True means the framework accepts responsibility for completing release.
# False or exceptions cause retries, which must be safe after partial release.
try:
released = self._kvcr._release_pin_callback(pin_handle)
except Exception:
logger.warning(
"KVCR release_pin failed for pin=%r", pin_handle, exc_info=True
)
if warn:
logger.warning(
"KVCR release_pin failed for pin=%r", pin_handle, exc_info=True
)
return False
if released is False:
logger.warning("KVCR release_pin failed for pin=%r", pin_handle)
if warn:
logger.warning("KVCR release_pin failed for pin=%r", pin_handle)
return False
return True

Expand Down
10 changes: 7 additions & 3 deletions tests/unit/_kvcr_test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -329,6 +329,10 @@ def initialize_xfer(
):
local_descs = list(local_descs)
remote_descs = list(remote_descs)
if len(local_descs) != len(remote_descs) or any(
local[1] != remote[1] for local, remote in zip(local_descs, remote_descs)
):
raise RuntimeError("NIXL rejected unaligned descriptors")
self.xfers.append(
(
op,
Expand All @@ -348,10 +352,10 @@ def transfer(self, handle):
handle - 1
]
if op == "WRITE" and remote_agent == self.name:
for local_index, remote_index in zip(local_indices, local_indices):
for local_index in local_indices:
src_addr, src_size, _ = local_descs[local_index]
dst_addr, dst_size, _ = remote_descs[remote_index]
ctypes.memmove(dst_addr, src_addr, min(src_size, dst_size))
dst_addr, _, _ = remote_descs[local_index]
ctypes.memmove(dst_addr, src_addr, src_size)
return "PROC"

def check_xfer_state(self, handle):
Expand Down
Loading
Loading