Skip to content
Closed
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
38 changes: 29 additions & 9 deletions modelexpress_client/python/modelexpress/nixl_transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,10 @@ def __init__(
self._metadata: bytes = b""
self._tensor_descriptors: list[TensorDescriptor] = []
self._tensors: dict[str, torch.Tensor] = {}
# Registration descriptors must be deregistered before destroying the
# UCX-backed NIXL agent. Dropping an agent with live GPU registrations
# can abort inside ucp_worker_destroy during framework teardown.
self._registered_memory: list[Any] = []
# Remote agents this manager has loaded, so shutdown can disconnect them.
# Maps agent name -> (ip, port) for agents reached over the P2P socket, or
# None for agents loaded from a metadata blob.
Expand Down Expand Up @@ -360,15 +364,19 @@ def register_tensors(
alloc_tuples = [
(base, size, self._device_id, "") for base, size in allocations
]
self._agent.register_memory(
alloc_tuples,
mem_type=self._accelerator_backend.nixl_mem_type,
backends=self._backends,
self._registered_memory.append(
self._agent.register_memory(
alloc_tuples,
mem_type=self._accelerator_backend.nixl_mem_type,
backends=self._backends,
)
)
reg_count = len(allocations)
else:
tensor_list = list(tensors.values())
self._agent.register_memory(tensor_list, backends=self._backends)
self._registered_memory.append(
self._agent.register_memory(tensor_list, backends=self._backends)
)
reg_count = len(tensor_list)
nixl_reg_time = time.perf_counter() - nixl_reg_start

Expand Down Expand Up @@ -507,10 +515,12 @@ def register_arena(
return self.register_tensors(tensors, force_per_tensor=True)

nixl_reg_start = time.perf_counter()
self._agent.register_memory(
[(base, used, self._device_id, "")],
mem_type=self._accelerator_backend.nixl_mem_type,
backends=self._backends,
self._registered_memory.append(
self._agent.register_memory(
[(base, used, self._device_id, "")],
mem_type=self._accelerator_backend.nixl_mem_type,
backends=self._backends,
)
)
nixl_reg_time = time.perf_counter() - nixl_reg_start

Expand Down Expand Up @@ -1221,6 +1231,16 @@ def shutdown(self) -> None:
atexit.unregister(self.shutdown)
self._atexit_registered = False
disconnected = self.disconnect_remote_agents()
if self._agent is not None:
for registered in reversed(self._registered_memory):
try:
self._agent.deregister_memory(registered)
except Exception:
logger.warning(
"Failed to deregister NIXL memory during shutdown",
exc_info=True,
)
self._registered_memory = []
self._agent = None
self._metadata = b""
self._tensor_descriptors = []
Expand Down
24 changes: 24 additions & 0 deletions modelexpress_client/python/tests/test_nixl_peer_lifecycle.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@ class FakeAgent:

def __init__(self, fail_remove: bool = False):
self.removed: list[str] = []
self.deregistered: list[object] = []
self.fail_remove = fail_remove
self._fetched: set[str] = set()

Expand All @@ -49,6 +50,9 @@ def remove_remote_agent(self, name: str):
raise RuntimeError(f"remote metadata for agent '{name}' not found")
self.removed.append(name)

def deregister_memory(self, registered):
self.deregistered.append(registered)


def _manager(agent=None, metadata=b"md", accelerator=None):
mgr = NixlTransferManager(
Expand Down Expand Up @@ -157,6 +161,26 @@ def spy(name):
assert seen["agent_alive"] is True
assert mgr._agent is None

def test_registered_memory_is_released_before_agent_is_dropped(self):
agent = FakeAgent()
mgr = _manager(agent=agent)
first, second = object(), object()
mgr._registered_memory = [first, second]
seen = []
real_deregister = agent.deregister_memory

def spy(registered):
seen.append((registered, mgr._agent is agent))
real_deregister(registered)

agent.deregister_memory = spy
mgr.shutdown()

assert seen == [(second, True), (first, True)]
assert agent.deregistered == [second, first]
assert mgr._registered_memory == []
assert mgr._agent is None

def test_removing_a_peer_twice_is_harmless(self):
"""The peer may already be gone, e.g. it sent us NIXLCOMM:INVL on exit."""
agent = FakeAgent()
Expand Down
Loading