diff --git a/modelexpress_client/python/modelexpress/nixl_transfer.py b/modelexpress_client/python/modelexpress/nixl_transfer.py index 4cb3aeeb..48028ca3 100644 --- a/modelexpress_client/python/modelexpress/nixl_transfer.py +++ b/modelexpress_client/python/modelexpress/nixl_transfer.py @@ -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. @@ -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 @@ -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 @@ -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 = [] diff --git a/modelexpress_client/python/tests/test_nixl_peer_lifecycle.py b/modelexpress_client/python/tests/test_nixl_peer_lifecycle.py index 35011739..16ad4781 100644 --- a/modelexpress_client/python/tests/test_nixl_peer_lifecycle.py +++ b/modelexpress_client/python/tests/test_nixl_peer_lifecycle.py @@ -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() @@ -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( @@ -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()