Skip to content
Open
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
1 change: 1 addition & 0 deletions lightllm/server/pd_io_struct.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,6 +170,7 @@ class PDChunckedTransTask:
start_trans_time: float = None # 用于标记传输开始的时间。同时标记是否正在传输中

error_info: Optional[str] = None
transfer_quiesced: bool = False
transfer_time_out_secs: int = 66
page_kind: str = "kv"
# Only valid for the local task owner; remote notify copies may carry the sender-local req_idx.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -138,6 +138,7 @@ def __init__(
self.recv_task_group_queue = queue.Queue()
self.waiting_dict_lock = threading.Lock()
self.waiting_dict: Dict[str, PDChunckedTransTask] = {}
self.quarantined_dict: Dict[str, PDChunckedTransTask] = {}
self.request_page_task_queue = queue.Queue()
self.ready_page_task_queue = queue.Queue()
self.success_queue = queue.Queue()
Expand Down Expand Up @@ -198,7 +199,7 @@ def _abort(self, request_id: int, error_info: str = "aborted req"):

for trans_task in aborted_tasks:
trans_task.error_info = error_info
self.failed_queue.put(trans_task)
self._queue_failed_task(trans_task)
return

@log_exception
Expand Down Expand Up @@ -263,15 +264,20 @@ def accept_peer_task_loop(

# 请求有错误
if notify_obj.error_info is not None:
# 直接清理掉所有的相关请求。
quarantined_task = None
with self.waiting_dict_lock:
local_trans_task = self.waiting_dict.pop(notify_obj.get_key(), None)
if notify_obj.transfer_quiesced:
quarantined_task = self.quarantined_dict.pop(notify_obj.get_key(), None)
if local_trans_task is not None:
local_trans_task.error_info = notify_obj.error_info
# 软性的调整超时时间,防止一些特殊情况,过快的释放task
# 占用的page 页面,导致多p 复写引起脏内容的问题。
local_trans_task.transfer_time_out_secs = 12
self.failed_queue.put(local_trans_task)
local_trans_task.transfer_quiesced = notify_obj.transfer_quiesced

if local_trans_task is not None:
self._queue_failed_task(local_trans_task)

if quarantined_task is not None:
self._recycle_quarantined_page(quarantined_task)

self._abort(
request_id=notify_obj.request_id,
Expand Down Expand Up @@ -311,13 +317,18 @@ def accept_peer_task_loop(

# prefill 写完数据到了 done 阶段
if remote_trans_task.write_stage == "done":
quarantined_task = None
with self.waiting_dict_lock:
local_trans_task = self.waiting_dict.pop(remote_trans_task.get_key(), None)
if local_trans_task is None:
quarantined_task = self.quarantined_dict.pop(remote_trans_task.get_key(), None)
if local_trans_task is not None:
local_trans_task.first_gen_token_id = remote_trans_task.first_gen_token_id
local_trans_task.first_gen_token_logprob = remote_trans_task.first_gen_token_logprob
self.ready_page_task_queue.put(local_trans_task)
logger.info(f"recv WRITE done from prefill: {remote_trans_task.to_str()}")
elif quarantined_task is not None:
self._recycle_quarantined_page(quarantined_task)
else:
# Same race as the WRITE request stage: decode may have cleaned the
# waiting task because the request was aborted, then a late done notify
Expand Down Expand Up @@ -346,7 +357,7 @@ def _check_tasks_time_out(self):

for trans_task in timeout_tasks:
trans_task.error_info = "time out in accept_peer_task_loop"
self.failed_queue.put(trans_task)
self._queue_failed_task(trans_task)
return

@log_exception
Expand All @@ -369,7 +380,7 @@ def request_page_loop(self):
logger.exception(str(e))
self.transporter.remove_remote_agent(peer_name=trans_task.prefill_agent_name)
trans_task.error_info = f"send write ready task to prefill node failed: {str(e)}"
self.failed_queue.put(trans_task)
self._queue_failed_task(trans_task)
continue

return
Expand Down Expand Up @@ -432,9 +443,10 @@ def fail_loop(self):
while True:
trans_task: PDChunckedTransTask = self.failed_queue.get()

# 回收页面
if trans_task.dst_page_index is not None:
self.page_index_queue.put(trans_task.dst_page_index)
if trans_task.transfer_quiesced:
self.page_index_queue.put(trans_task.dst_page_index)
trans_task.dst_page_index = None

if trans_task.xfer_handle is not None:
self.transporter.release_xfer_handle(trans_task.xfer_handle)
Expand All @@ -449,4 +461,18 @@ def fail_loop(self):
request_id=trans_task.request_id,
error_info=trans_task.error_info,
)
self.transporter.send_error_info_to_prefill_node(trans_task=trans_task)
if not trans_task.transfer_quiesced:
self.transporter.send_error_info_to_prefill_node(trans_task=trans_task)

def _recycle_quarantined_page(self, trans_task: PDChunckedTransTask):
if trans_task.dst_page_index is not None:
self.page_index_queue.put(trans_task.dst_page_index)
trans_task.dst_page_index = None
trans_task.transfer_quiesced = True
logger.info(f"recycle quiesced decode page for failed task: {trans_task.to_str()}")

def _queue_failed_task(self, trans_task: PDChunckedTransTask):
if trans_task.dst_page_index is not None and not trans_task.transfer_quiesced:
with self.waiting_dict_lock:
self.quarantined_dict[trans_task.get_key()] = trans_task
self.failed_queue.put(trans_task)
Original file line number Diff line number Diff line change
Expand Up @@ -183,7 +183,7 @@ def send_error_info_to_decode_node(self, trans_task: PDChunckedTransTask):
new_trans_task.prefill_num_pages = self.num_pages
new_trans_task.prefill_page_reg_desc = self.local_page_mem_desc
self._send_task_notif(trans_task.decode_agent_name, new_trans_task)
return
return True

def write_blocks_paged(self, trans_task: PDChunckedTransTask) -> "_NcclXferHandle":
assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -229,11 +229,12 @@ def send_error_info_to_decode_node(self, trans_task: PDChunckedTransTask):
remote_agent_name=decode_agent_name,
notif_msg=pickle.dumps(new_trans_task),
)
return True
except BaseException as e:
logger.error(f"send error info to decode node failed: {trans_task.to_str()}")
logger.exception(str(e))
self.remove_remote_agent(peer_name=decode_agent_name)
return
return False

def write_blocks_paged(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,7 @@ def __init__(
self.write_peer_kv_queue = queue.Queue()
self.success_queue = queue.Queue()
self.failed_queue = queue.Queue()
self.draining_queue = queue.Queue()

self.page_index_queue = queue.Queue()
for page_index in range(self.args.pd_kv_page_num):
Expand All @@ -133,6 +134,7 @@ def __init__(
self.update_task_status_loop,
self.success_loop,
self.fail_loop,
self.drain_failed_xfer_loop,
]:
threading.Thread(target=func, daemon=True).start()
return
Expand Down Expand Up @@ -377,11 +379,12 @@ def success_loop(self):
torch.cuda.set_device(self.device_id)
while True:
trans_task: PDChunckedTransTask = self.success_queue.get()
# 写回后,回收页面
if trans_task.src_page_index is not None:
self.page_index_queue.put(trans_task.src_page_index)
if trans_task.xfer_handle is not None:
self.transporter.release_xfer_handle(trans_task.xfer_handle)
trans_task.xfer_handle = None
if trans_task.src_page_index is not None:
self.page_index_queue.put(trans_task.src_page_index)
trans_task.src_page_index = None

ret = trans_task.createRetObj()
ret.first_gen_token_id = None
Expand All @@ -399,16 +402,69 @@ def fail_loop(self):
while True:
trans_task: PDChunckedTransTask = self.failed_queue.get()

# 回收页面
if trans_task.src_page_index is not None:
self.page_index_queue.put(trans_task.src_page_index)
if trans_task.xfer_handle is not None:
self.transporter.release_xfer_handle(trans_task.xfer_handle)
if trans_task.xfer_handle is None:
if trans_task.src_page_index is not None:
self.page_index_queue.put(trans_task.src_page_index)
trans_task.src_page_index = None
trans_task.transfer_quiesced = True
if trans_task.error_info is not None:
self.draining_queue.put(trans_task)

ret = trans_task.createRetObj()
self.task_out_queue.put(ret)
logger.info(f"trans task ret fail:{ret}")

if trans_task.error_info is not None:
self._abort(request_id=trans_task.request_id, error_info=trans_task.error_info)
self.transporter.send_error_info_to_decode_node(trans_task=trans_task)

@log_exception
def drain_failed_xfer_loop(self):
torch.cuda.set_device(self.device_id)
draining_tasks: Dict[str, PDChunckedTransTask] = {}
while True:
try:
trans_task = self.draining_queue.get(timeout=0.001)
draining_tasks[trans_task.get_key()] = trans_task
except queue.Empty:
pass

for key, trans_task in list(draining_tasks.items()):
try:
status = self._try_finish_failed_xfer_drain(trans_task)
if status is None:
continue
except BaseException as e:
if trans_task.transfer_quiesced:
logger.error(
f"failed to send quiesced ack, decode page remains quarantined: {trans_task.to_str()}"
)
else:
logger.error(f"failed to drain xfer task, keeping page quarantined: {trans_task.to_str()}")
logger.exception(str(e))
continue

draining_tasks.pop(key, None)
logger.info(f"failed xfer task reached terminal state {status}: {trans_task.to_str()}")

def _try_finish_failed_xfer_drain(self, trans_task: PDChunckedTransTask) -> Optional[str]:
if trans_task.transfer_quiesced:
sent = self.transporter.send_error_info_to_decode_node(trans_task=trans_task)
if sent is False:
raise RuntimeError("failed to send quiesced transfer acknowledgement")
return "QUIESCED"

status = self.transporter.check_task_status(trans_task=trans_task)
if status not in ["DONE", "ERR"]:
return None

self.transporter.release_xfer_handle(trans_task.xfer_handle)
trans_task.xfer_handle = None
if trans_task.src_page_index is not None:
self.page_index_queue.put(trans_task.src_page_index)
trans_task.src_page_index = None

trans_task.transfer_quiesced = True
sent = self.transporter.send_error_info_to_decode_node(trans_task=trans_task)
if sent is False:
raise RuntimeError("failed to send quiesced transfer acknowledgement")
return status
Loading
Loading