Skip to content

Commit 60cd914

Browse files
committed
fix
1 parent 2f98173 commit 60cd914

1 file changed

Lines changed: 150 additions & 80 deletions

File tree

lightllm/server/router/model_infer/mode_backend/pd_nixl/prefill_node_impl/prefill_trans_process.py

Lines changed: 150 additions & 80 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,12 @@
44
import threading
55
import setproctitle
66
import torch.multiprocessing as mp
7-
import collections
87
import queue
98
import pickle
10-
from typing import List, Dict, Union, Deque, Optional
9+
from typing import List, Dict, Optional
1110
from lightllm.utils.log_utils import init_logger
1211
from 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
1513
from lightllm.utils.graceful_utils import graceful_registry
1614
from lightllm.server.core.objs import StartArgs
1715
from ..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

Comments
 (0)