diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index bd5e11f05..8bcd23abb 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -1,6 +1,7 @@ import pickle import copy import os +import threading import time from dataclasses import dataclass from typing import Dict @@ -40,6 +41,10 @@ def __init__(self, node_id: int, tp_idx: int, kv_move_buffer: Tensor): self.nixl_agent = NixlWrapper(self.agent_name, conf) self._register_kv_move_buffer(kv_move_buffer=kv_move_buffer) self.remote_agents: Dict[str, PDAgentMetadata] = {} + # remote_agents is read and mutated from several worker threads (recv, dispatch, + # accept_peer, request/ready page loops, ...), so every access needs to go through + # this lock to avoid concurrent add/remove races on the same peer (see GH-1470). + self._remote_agents_lock = threading.Lock() return @property @@ -73,36 +78,38 @@ def _create_paged_xfer_handles(self, reg_desc: "nixlBind.nixlRegDList", page_num return self.nixl_agent.prep_xfer_dlist(agent_name, descs, "VRAM") def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): - if remote_agent.agent_name in self.remote_agents: - return + with self._remote_agents_lock: + if remote_agent.agent_name in self.remote_agents: + return - start_time = time.time() + start_time = time.time() - peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) - if isinstance(peer_name, bytes): - peer_name = peer_name.decode() + peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) + if isinstance(peer_name, bytes): + peer_name = peer_name.decode() - assert ( - peer_name == remote_agent.agent_name - ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" + assert ( + peer_name == remote_agent.agent_name + ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" - page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) - kv_page_xfer_handles = self._create_paged_xfer_handles( - page_mem_desc, remote_agent.num_pages, agent_name=peer_name - ) - remote_agent.page_xfer_handles = kv_page_xfer_handles + page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) + kv_page_xfer_handles = self._create_paged_xfer_handles( + page_mem_desc, remote_agent.num_pages, agent_name=peer_name + ) + remote_agent.page_xfer_handles = kv_page_xfer_handles - logger.info( - f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" - ) + logger.info( + f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" + ) - self.remote_agents[remote_agent.agent_name] = remote_agent + self.remote_agents[remote_agent.agent_name] = remote_agent return def remove_remote_agent(self, peer_name: str): - if peer_name in self.remote_agents: + with self._remote_agents_lock: + remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) + if remote_agent is not None: try: - remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) assert remote_agent.agent_name == peer_name self.nixl_agent.remove_remote_agent(remote_agent.agent_name) if remote_agent.page_xfer_handles is not None: