44import threading
55import setproctitle
66import torch .multiprocessing as mp
7- import collections
87import queue
98import pickle
10- from typing import List , Dict , Union , Deque , Optional
9+ from typing import List , Dict , Optional
1110from lightllm .utils .log_utils import init_logger
1211from lightllm .common .kv_cache_mem_manager import MemoryManager
13- from lightllm .server .pd_io_struct import NIXLChunckedTransTask , NIXLChunckedTransTaskRet
14- from lightllm .utils .device_utils import kv_trans_use_p2p
12+ from lightllm .server .pd_io_struct import NIXLChunckedTransTask
1513from lightllm .utils .graceful_utils import graceful_registry
1614from lightllm .server .core .objs import StartArgs
1715from ..nixl_kv_transporter import NixlKVTransporter
@@ -45,9 +43,8 @@ def _init_env(
4543
4644 import os
4745
48- # prefill 节点不一定需要 mps 来协调,所以优先级设置为 1.
49- # 本身并不产生严重的阻塞。
50- os .environ ["CUDA_MPS_CLIENT_PRIORITY" ] = "1"
46+ # prefill source-side page copy and UCX progress are on the request critical path.
47+ os .environ ["CUDA_MPS_CLIENT_PRIORITY" ] = "0"
5148
5249 torch .backends .cudnn .enabled = False
5350 setproctitle .setproctitle (f"lightllm::{ get_unique_server_name ()} ::nixl_prefill_trans:Device{ device_id } " )
@@ -103,15 +100,18 @@ def __init__(
103100 kv_move_buffer = cur_mem_manager .alloc_paged_kv_move_buffer (
104101 page_num = self .args .nixl_pd_kv_page_num , page_size = self .args .nixl_pd_kv_page_size
105102 )
106- self .copy_cuda_stream = torch .cuda .Stream ()
103+ self .copy_cuda_stream = torch .cuda .Stream (priority = - 1 )
107104 self .transporter = NixlKVTransporter (
108105 node_id = self .args .pd_node_id , tp_idx = device_id , kv_move_buffer = kv_move_buffer
109106 )
110107 self .waiting_dict_lock = threading .Lock ()
111108 self .waiting_dict : Dict [str , NIXLChunckedTransTask ] = {}
109+ self .waiting_nixl_write_task_lock = threading .Lock ()
110+ self .waiting_nixl_write_task_dict : Dict [str , NIXLChunckedTransTask ] = {}
112111
113112 self .local_copy_kv_queue = queue .Queue ()
114- self .notify_peer_read_kv_queue = queue .Queue ()
113+ self .ready_transfer_queue = queue .Queue ()
114+ self .write_peer_kv_queue = queue .Queue ()
115115 self .success_queue = queue .Queue ()
116116 self .failed_queue = queue .Queue ()
117117
@@ -125,14 +125,30 @@ def __init__(
125125 for func in [
126126 self .recv_task_loop ,
127127 self .local_copy_kv_loop ,
128- self .notify_peer_to_read_kv_loop ,
128+ self .ready_transfer_loop ,
129+ self .accept_decode_write_task_loop ,
130+ self .write_peer_kv_loop ,
129131 self .update_task_status_loop ,
130132 self .success_loop ,
131133 self .fail_loop ,
132134 ]:
133135 threading .Thread (target = func , daemon = True ).start ()
134136 return
135137
138+ def _warmup (self ):
139+ for dp_index in range (self .args .dp // self .args .nnodes ):
140+ with torch .cuda .stream (stream = self .copy_cuda_stream ):
141+ cur_mem = self .mem_managers [self .device_id ]
142+ cur_mem .write_mem_to_page_kv_move_buffer (
143+ mem_indexes = [0 ],
144+ page_index = 0 ,
145+ dp_index = dp_index ,
146+ mem_managers = self .mem_managers ,
147+ dp_world_size = self .dp_world_size ,
148+ )
149+ torch .cuda .current_stream ().synchronize ()
150+ return
151+
136152 @log_exception
137153 def recv_task_loop (self ):
138154 torch .cuda .set_device (self .device_id )
@@ -168,65 +184,26 @@ def local_copy_kv_loop(self):
168184 sync_event = torch .cuda .Event ()
169185 sync_event .record ()
170186
171- self .notify_peer_read_kv_queue .put ((sync_event , trans_task ))
172- return
173-
174- def _warmup (self ):
175- for dp_index in range (self .args .dp // self .args .nnodes ):
176- with torch .cuda .stream (stream = self .copy_cuda_stream ):
177- cur_mem = self .mem_managers [self .device_id ]
178- cur_mem .write_mem_to_page_kv_move_buffer (
179- mem_indexes = [0 ],
180- page_index = 0 ,
181- dp_index = dp_index ,
182- mem_managers = self .mem_managers ,
183- dp_world_size = self .dp_world_size ,
184- )
185- torch .cuda .current_stream ().synchronize ()
187+ self .ready_transfer_queue .put ((sync_event , trans_task ))
186188 return
187189
188190 @log_exception
189- def notify_peer_to_read_kv_loop (self ):
191+ def ready_transfer_loop (self ):
190192 torch .cuda .set_device (self .device_id )
191193 while True :
192- sync_event , trans_task = self .notify_peer_read_kv_queue .get ()
194+ sync_event , trans_task = self .ready_transfer_queue .get ()
193195 trans_task : NIXLChunckedTransTask = trans_task
194196 sync_event : torch .cuda .Event = sync_event
195-
196197 sync_event .synchronize ()
197-
198- trans_task . start_trans_time = time . time ( )
198+ self . transporter . send_write_request_task_to_decode_node ( trans_task )
199+ logger . info ( f"send WRITE request to decode: { trans_task . get_key () } " )
199200 with self .waiting_dict_lock :
200201 self .waiting_dict [trans_task .get_key ()] = trans_task
201-
202- try :
203- self .transporter .send_readtask_to_decode_node (trans_task = trans_task )
204- except BaseException as e :
205- logger .error (f"send readtask to decode node failed: { trans_task .to_str ()} " )
206- logger .exception (str (e ))
207- self .transporter .remove_remote_agent (peer_name = trans_task .decode_agent_name )
208-
209- with self .waiting_dict_lock :
210- trans_task = self .waiting_dict .pop (trans_task .get_key (), None )
211-
212- if trans_task is not None :
213- trans_task .error_info = f"send readtask to decode node failed: { str (e )} "
214- self .failed_queue .put (trans_task )
215- continue
216-
217- logger .info (f"send readtask to decode: { trans_task .to_str ()} " )
218202 return
219203
220204 @log_exception
221- def update_task_status_loop (
222- self ,
223- ):
205+ def accept_decode_write_task_loop (self ):
224206 while True :
225- if len (self .waiting_dict ) == 0 :
226- time .sleep (0.001 )
227- continue
228-
229- # notify update
230207 try :
231208 notifies_dict = self .transporter .get_new_notifs ()
232209 except BaseException as e :
@@ -239,44 +216,133 @@ def update_task_status_loop(
239216 for notify in _notify_list :
240217 try :
241218 notify_obj = pickle .loads (notify )
242- except :
219+ except BaseException :
243220 notify_obj = None
244221
245- if isinstance (notify_obj , NIXLChunckedTransTaskRet ):
246- key = notify_obj .get_key ()
247- with self .waiting_dict_lock :
248- trans_task = self .waiting_dict .pop (key , None )
249-
250- if trans_task is not None :
251- trans_task .error_info = notify_obj .error_info
252- if trans_task .error_info is not None :
253- self .failed_queue .put (trans_task )
254- else :
255- self .success_queue .put (trans_task )
256- else :
257- logger .warning (f"can not find trans task for ret: { notify_obj } " )
258-
259- # check time_out update
222+ if not isinstance (notify_obj , NIXLChunckedTransTask ):
223+ continue
224+
225+ if notify_obj .nixl_write_stage != "ready" :
226+ logger .warning (f"ignore unknown prefill WRITE notify stage: { notify_obj .to_str ()} " )
227+ continue
228+
229+ key = notify_obj .get_key ()
230+ with self .waiting_dict_lock :
231+ trans_task = self .waiting_dict .pop (key , None )
232+
233+ if trans_task is None :
234+ logger .warning (
235+ f"can not find pending WRITE request for ready notify: { notify_obj .to_str ()} "
236+ )
237+ continue
238+
239+ trans_task .nixl_dst_page_index = notify_obj .nixl_dst_page_index
240+ logger .info (
241+ f"recv WRITE ready from decode request_id={ trans_task .request_id } "
242+ f"kv=[{ trans_task .start_kv_index } ,{ trans_task .end_kv_index } ) "
243+ f"src_page={ trans_task .nixl_src_page_index } dst_page={ trans_task .nixl_dst_page_index } "
244+ )
245+ self .write_peer_kv_queue .put (trans_task )
246+
260247 self ._check_tasks_time_out ()
261248
249+ if not notifies_dict :
250+ time .sleep (0.001 )
251+ return
252+
262253 def _check_tasks_time_out (self ):
263254 with self .waiting_dict_lock :
264- keys = list (self .waiting_dict .keys ())
255+ timeout_tasks = []
256+ for key , trans_task in list (self .waiting_dict .items ()):
257+ if trans_task .time_out ():
258+ timeout_tasks .append (self .waiting_dict .pop (key ))
259+
260+ for trans_task in timeout_tasks :
261+ trans_task .error_info = "time out waiting decode WRITE ready"
262+ self .failed_queue .put (trans_task )
263+ return
265264
266- for key in keys :
267- with self .waiting_dict_lock :
268- trans_task = self .waiting_dict .pop (key , None )
265+ @log_exception
266+ def write_peer_kv_loop (self ):
267+ torch .cuda .set_device (self .device_id )
268+ while True :
269+ trans_task = self .write_peer_kv_queue .get ()
270+ trans_task : NIXLChunckedTransTask = trans_task
269271
270- if trans_task is not None and trans_task .time_out ():
271- trans_task .error_info = "time out in update_task_status_loop"
272+ try :
273+ xfer_handle = self .transporter .write_blocks_paged (trans_task = trans_task )
274+ trans_task .xfer_handle = xfer_handle
275+ except BaseException as e :
276+ logger .error (f"write_blocks_paged failed: { trans_task .to_str ()} " )
277+ logger .exception (str (e ))
278+ self .transporter .remove_remote_agent (peer_name = trans_task .decode_agent_name )
279+
280+ trans_task .error_info = f"write_blocks_paged failed: { str (e )} "
272281 self .failed_queue .put (trans_task )
273282 continue
274283
275- if trans_task is not None :
276- with self .waiting_dict_lock :
277- self .waiting_dict [trans_task .get_key ()] = trans_task
284+ trans_task .start_trans_time = time .time ()
285+ with self .waiting_nixl_write_task_lock :
286+ self .waiting_nixl_write_task_dict [trans_task .get_key ()] = trans_task
287+ logger .info (f"start WRITE to decode node: { trans_task .to_str ()} " )
278288 return
279289
290+ @log_exception
291+ def update_task_status_loop (
292+ self ,
293+ ):
294+ while True :
295+ if len (self .waiting_nixl_write_task_dict ) == 0 :
296+ time .sleep (0.001 )
297+ continue
298+
299+ with self .waiting_nixl_write_task_lock :
300+ tasks = list (self .waiting_nixl_write_task_dict .values ())
301+
302+ for trans_task in tasks :
303+ ret = self .transporter .check_task_status (trans_task = trans_task )
304+ if ret == "DONE" :
305+ with self .waiting_nixl_write_task_lock :
306+ trans_task = self .waiting_nixl_write_task_dict .pop (trans_task .get_key (), None )
307+ if trans_task is None :
308+ continue
309+ if self .transporter .capture_telemetry :
310+ telem = self .transporter .nixl_agent .get_xfer_telemetry (trans_task .xfer_handle )
311+ total_us = telem .xferDuration
312+ post_us = telem .postDuration
313+ backend_us = telem .xferDuration - telem .postDuration
314+ nixl_backend = self .transporter .nixl_agent .query_xfer_backend (trans_task .xfer_handle )
315+ logger .info (
316+ f"write trans task request_id={ trans_task .request_id } "
317+ f"kv=[{ trans_task .start_kv_index } ,{ trans_task .end_kv_index } ) "
318+ f"src_page={ trans_task .nixl_src_page_index } dst_page={ trans_task .nixl_dst_page_index } "
319+ f"xfer time: { total_us :.3f} us, "
320+ f"post time: { post_us :.3f} us, backend time: { backend_us :.3f} us, "
321+ f"nixl_backend: { nixl_backend } , total_bytes: { telem .totalBytes } "
322+ )
323+ self .transporter .send_write_done_task_to_decode_node (trans_task )
324+ logger .info (
325+ f"send WRITE done nixl notify "
326+ f"request_id={ trans_task .request_id } "
327+ f"kv=[{ trans_task .start_kv_index } ,{ trans_task .end_kv_index } ) "
328+ f"src_page={ trans_task .nixl_src_page_index } dst_page={ trans_task .nixl_dst_page_index } "
329+ )
330+ self .success_queue .put (trans_task )
331+ elif ret == "ERR" :
332+ with self .waiting_nixl_write_task_lock :
333+ trans_task = self .waiting_nixl_write_task_dict .pop (trans_task .get_key (), None )
334+ if trans_task is not None :
335+ trans_task .error_info = "xfer error"
336+ self .failed_queue .put (trans_task )
337+ elif trans_task .time_out ():
338+ with self .waiting_nixl_write_task_lock :
339+ trans_task = self .waiting_nixl_write_task_dict .pop (trans_task .get_key (), None )
340+ if trans_task is not None :
341+ trans_task .error_info = "time out in update_task_status_loop"
342+ self .failed_queue .put (trans_task )
343+
344+ time .sleep (0.001 )
345+
280346 @log_exception
281347 def success_loop (self ):
282348 torch .cuda .set_device (self .device_id )
@@ -285,6 +351,8 @@ def success_loop(self):
285351 # 写回后,回收页面
286352 if trans_task .nixl_src_page_index is not None :
287353 self .page_index_queue .put (trans_task .nixl_src_page_index )
354+ if trans_task .xfer_handle is not None :
355+ self .transporter .release_xfer_handle (trans_task .xfer_handle )
288356
289357 ret = trans_task .createRetObj ()
290358 ret .first_gen_token_id = None
@@ -301,6 +369,8 @@ def fail_loop(self):
301369 # 回收页面
302370 if trans_task .nixl_src_page_index is not None :
303371 self .page_index_queue .put (trans_task .nixl_src_page_index )
372+ if trans_task .xfer_handle is not None :
373+ self .transporter .release_xfer_handle (trans_task .xfer_handle )
304374
305375 ret = trans_task .createRetObj ()
306376 self .task_out_queue .put (ret )
0 commit comments