Skip to content

Commit 9de9265

Browse files
committed
prepare request in prefill instance by multi threads
1 parent 0ec9625 commit 9de9265

5 files changed

Lines changed: 307 additions & 323 deletions

File tree

fastdeploy/cache_manager/cache_messager.py

Lines changed: 26 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -613,12 +613,16 @@ def __init__(
613613
)
614614

615615
self.gpu_id = gpu_id
616-
self.cache_info = dict()
616+
self.cache_info = dict() # {'request_id': cache_info_dict}
617617
self.rank_id = self.rank + local_data_parallel_id * self.nranks
618618
self.engine_cache_task_thread_lock = threading.Lock()
619-
self.engine_cache_tasks = [dict() for _ in range(512)]
620-
self.idx_cache_task_dict = {}
621-
self.cache_prefilled_engine_ids_queue = queue.Queue() # keep batch slot index for each prefill step
619+
self.engine_cache_tasks = [
620+
dict() for _ in range(512)
621+
] # {'layer_id': {'prefilled_layer_idx': xx, 'prefilled_block_num': xx}}
622+
self.idx_cache_task_dict = {} # {'solt_idx': cache_info_dict}
623+
self.cache_prefilled_engine_ids_queue = (
624+
queue.Queue()
625+
) # [(solt_idx1, prefilled_token_num1), (solt_idx2, prefilled_token_num2)]
622626
if splitwise_role == "prefill":
623627
consume_signals_thread = threading.Thread(target=self.consume_signals)
624628
consume_signals_thread.daemon = True
@@ -638,7 +642,6 @@ def _add_cache_task_thread(self):
638642
while True:
639643
try:
640644
cache_info = self.engine_worker_queue.get_cache_info()
641-
finished_add_cache_task_req_ids = []
642645
if cache_info:
643646
logger.debug(f"Get cache info from engine worker queue, {cache_info}")
644647
self.engine_worker_queue.cache_info_barrier.wait()
@@ -647,7 +650,6 @@ def _add_cache_task_thread(self):
647650
self.cache_info[info["request_id"]].update(info)
648651
current_info = self.cache_info[info["request_id"]]
649652
assert "dest_block_ids" in current_info and "src_block_ids" in current_info
650-
finished_add_cache_task_req_ids.append(info["request_id"])
651653
decode_cached_block_num = len(current_info["src_block_ids"]) - len(
652654
current_info["dest_block_ids"]
653655
)
@@ -659,17 +661,13 @@ def _add_cache_task_thread(self):
659661
current_info["sended_layer_id"] = -1
660662
current_info["sended_block_num"] = current_info["decode_cached_tokens"] // self.block_size
661663
current_info["status"] = "init"
662-
logger.info(f"Get cache info from D: finish add cache task: {current_info}")
664+
logger.info(f"Get cache info and finish add cache task: {current_info}")
663665
self.cache_info[info["request_id"]] = current_info
664666
self.idx_cache_task_dict[current_info["current_id"]] = current_info
665667
else:
666-
logger.info(f"Get cache info from P: {info}")
668+
logger.info(f"Get cache info: {info}")
667669
self.cache_info[info["request_id"]] = info
668670

669-
if finished_add_cache_task_req_ids:
670-
logger.info(f"Put processed tasks into engine worker queue: {finished_add_cache_task_req_ids}")
671-
self.engine_worker_queue.put_finished_add_cache_task_req(finished_add_cache_task_req_ids)
672-
self.engine_worker_queue.finish_add_cache_task_barrier.wait()
673671
else:
674672
time.sleep(0.001)
675673
except Exception as e:
@@ -687,10 +685,12 @@ def prefill_layerwise_send_cache_thread(self):
687685
block_start_end_list = []
688686
current_prefilled_token_num_list = []
689687
for engine_index, current_step_prefilled_token_num in batch_engine_signals:
688+
self._maybe_wait_for_cache_task(engine_index)
690689
assert (
691690
engine_index in self.idx_cache_task_dict
692691
), f"engine_index {engine_index} not in self.idx_cache_task_dict {self.idx_cache_task_dict}"
693692
block_id_start = self.idx_cache_task_dict[engine_index]["sended_block_num"]
693+
694694
prefilled_token_num = current_step_prefilled_token_num
695695
if (
696696
prefilled_token_num == self.idx_cache_task_dict[engine_index]["need_prefill_tokens"]
@@ -917,6 +917,20 @@ def _handle_connect_task(self):
917917
except Exception as e:
918918
logger.error(f"handle_connect_task has exception: {e}, {traceback.format_exc()}")
919919

920+
def _maybe_wait_for_cache_task(self, engine_index):
921+
# If cache messager does not get cache task from engine, just hang here for now
922+
wait_step = 1
923+
sleep_seconds = 0.005
924+
925+
while engine_index not in self.idx_cache_task_dict:
926+
time.sleep(sleep_seconds)
927+
wait_step += 1
928+
929+
if wait_step % 400 == 0:
930+
logger.warning(
931+
f"waiting cache task for engine_index: {engine_index}, cost_time: {wait_step * 0.005:.2f} s"
932+
)
933+
920934

921935
def main():
922936
device = args.device_id

fastdeploy/engine/common_engine.py

Lines changed: 13 additions & 205 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,6 @@
3030
import time
3131
import traceback
3232
import weakref
33-
from concurrent.futures import ThreadPoolExecutor
3433
from pathlib import Path
3534
from typing import Dict, List, Optional, Tuple
3635

@@ -42,6 +41,7 @@
4241
import fastdeploy.metrics.trace as tracing
4342
from fastdeploy.cache_manager.cache_data import CacheStatus
4443
from fastdeploy.config import FDConfig
44+
from fastdeploy.engine.common_engine_prepare_mixin import EngineServicePrepareMixin
4545
from fastdeploy.engine.register_manager import RegisterManager
4646
from fastdeploy.engine.request import (
4747
CompletionOutput,
@@ -115,7 +115,7 @@ def _format_worker_launch_failure_message(log_dir: str) -> str:
115115
return message
116116

117117

118-
class EngineService:
118+
class EngineService(EngineServicePrepareMixin):
119119
"""
120120
Base class containing common engine functionality
121121
"""
@@ -251,12 +251,13 @@ def start(self, async_llm_pid=None):
251251
self.start_worker_service(async_llm_pid)
252252

253253
if envs.ENABLE_V1_KVCACHE_SCHEDULER:
254-
self.insert_task_to_worker_thread = threading.Thread(
255-
target=self._schedule_request_to_worker_v1, daemon=True
256-
)
254+
self.prepare_request_thread = threading.Thread(target=self._prepare_request_v1, daemon=True)
255+
self.prepare_request_thread.start()
256+
self.schedule_request_thread = threading.Thread(target=self._schedule_request_to_worker_v1, daemon=True)
257+
self.schedule_request_thread.start()
257258
else:
258-
self.insert_task_to_worker_thread = threading.Thread(target=self._schedule_request_to_worker, daemon=True)
259-
self.insert_task_to_worker_thread.start()
259+
self.schedule_request_thread = threading.Thread(target=self._schedule_request_to_worker, daemon=True)
260+
self.schedule_request_thread.start()
260261
self.token_processor.tasks_queue = self.engine_worker_queue
261262
self.token_processor.run()
262263
if self.cfg.scheduler_config.splitwise_role == "decode":
@@ -879,215 +880,19 @@ def _schedule_request_to_worker_v1(self):
879880
Insert tasks to worker with scheduler v1 (ENABLE_V1_KVCACHE_SCHEDULER=1).
880881
"""
881882
tracing.trace_set_thread_info("Scheduler Task to Work")
882-
get_request_pool = ThreadPoolExecutor(max_workers=1)
883-
is_fetching = False
884-
885-
def _fetch_request():
886-
try:
887-
with self._pause_cond:
888-
self._pause_cond.wait_for(lambda: not self.is_paused)
889-
nonlocal is_fetching
890-
num_prefill_batch = min(
891-
int(self.resource_manager.available_batch()),
892-
self.cfg.max_prefill_batch,
893-
)
894-
895-
if self.cfg.scheduler_config.splitwise_role != "mixed":
896-
max_num_batched_tokens = self.cfg.scheduler_config.max_num_batched_tokens
897-
else:
898-
max_num_batched_tokens = self.cfg.model_config.max_model_len
899-
900-
available_blocks = self.cfg.cache_config.max_block_num_per_seq
901-
tasks = self.scheduler.get_requests(
902-
available_blocks=available_blocks,
903-
block_size=self.cfg.cache_config.block_size,
904-
reserved_output_blocks=0, # self.cfg.cache_config.enc_dec_block_num
905-
max_num_batched_tokens=max_num_batched_tokens,
906-
batch=num_prefill_batch,
907-
)
908-
for task in tasks:
909-
task.metrics.engine_get_req_time = time.time()
910-
trace_print(LoggingEventName.REQUEST_QUEUE_END, task.request_id, getattr(task, "user", ""))
911-
912-
if self.cfg.scheduler_config.splitwise_role == "decode":
913-
# TODO: refine scheduler to remove this limitation
914-
# Decode will process and schedule the request sent by prefill to engine,
915-
# so the same request sent by the decode api server will be ignored
916-
is_fetching = False
917-
return
918-
919-
if tasks:
920-
self.llm_logger.debug(
921-
f"Engine has fetched tasks from {self.scheduler.__class__.__name__}: {[task.request_id for task in tasks]}"
922-
)
923-
924-
if self.cfg.scheduler_config.splitwise_role == "prefill":
925-
for task in tasks:
926-
# start async preprocess
927-
self.resource_manager.apply_async_preprocess(task)
928-
need_delete_tasks = []
929-
if envs.PREFILL_CONTINUOUS_REQUEST_DECODE_RESOURCES:
930-
for task in tasks:
931-
# assure can allocate block ids in P
932-
while not self.resource_manager.preallocate_resource_in_p(task):
933-
time.sleep(0.005)
934-
self.llm_logger.debug(
935-
f"P has allocated resources and then ask D resource for request: {task.request_id}"
936-
)
937-
trace_print(
938-
LoggingEventName.ASK_DECODE_RESOURCE_START, task.request_id, getattr(task, "user", "")
939-
)
940-
task.metrics.ask_decode_resource_start_time = time.time()
941-
while True:
942-
self.split_connector.send_splitwise_tasks([task], task.idx)
943-
status, msg = self.split_connector.check_decode_allocated(task)
944-
if not status:
945-
self.llm_logger.warning(
946-
f"D failed to allocate resource for request {task.request_id}, try again."
947-
)
948-
time.sleep(0.05)
949-
else:
950-
task.metrics.ask_decode_resource_finish_time = time.time()
951-
trace_print(
952-
LoggingEventName.ASK_DECODE_RESOURCE_END,
953-
task.request_id,
954-
getattr(task, "user", ""),
955-
)
956-
break
957-
self.llm_logger.debug(f"D has allocated resource for request: {task.request_id}")
958-
else:
959-
for task in tasks:
960-
# assure can allocate block ids in P
961-
while not self.resource_manager.preallocate_resource_in_p(task):
962-
time.sleep(0.005)
963-
964-
self.llm_logger.debug(
965-
f"P has allocated resources and then ask D resource for req_id: {task.request_id}"
966-
)
967-
trace_print(
968-
LoggingEventName.ASK_DECODE_RESOURCE_START, task.request_id, getattr(task, "user", "")
969-
)
970-
task.metrics.ask_decode_resource_start_time = time.time()
971-
self.split_connector.send_splitwise_tasks([task], task.idx)
972-
973-
for task in tasks:
974-
# assure fetch block ids from D
975-
status, msg = self.split_connector.check_decode_allocated(task)
976-
task.metrics.ask_decode_resource_finish_time = time.time()
977-
trace_print(
978-
LoggingEventName.ASK_DECODE_RESOURCE_END, task.request_id, getattr(task, "user", "")
979-
)
980-
if not status:
981-
error_msg = (
982-
f"PD Error: prefill failed to apply for resource from decode, "
983-
f"req: {task.request_id}, msg:{msg}."
984-
)
985-
self.llm_logger.error(error_msg)
986-
self.scheduler.put_results(
987-
[
988-
RequestOutput(
989-
request_id=task.request_id,
990-
finished=True,
991-
error_code=500,
992-
error_msg=error_msg,
993-
)
994-
]
995-
)
996-
main_process_metrics.reschedule_req_num.inc()
997-
need_delete_tasks.append(task)
998-
continue
999-
for tmp_task in need_delete_tasks:
1000-
tasks.remove(tmp_task)
1001-
# release resource in P
1002-
self.resource_manager.pre_recycle_resource(tmp_task.request_id)
1003-
1004-
# to send cache info to cache messager
1005-
if tasks:
1006-
need_check_req_ids = [task.request_id for task in tasks]
1007-
self.split_connector.send_cache_info_to_messager(tasks, 0)
1008-
# ensure cache tasks has sent to cache_messager
1009-
need_check_req_ids = [task.request_id for task in tasks]
1010-
finished_ids, delete_tasks_list = [], []
1011-
while need_check_req_ids:
1012-
finished_ids.extend(self.engine_worker_queue.get_finished_add_cache_task_req())
1013-
self.llm_logger.debug(
1014-
f"P has successfully sent cache infos to cache messager for requests: {finished_ids}"
1015-
)
1016-
if finished_ids:
1017-
for task in tasks:
1018-
result = self.resource_manager.waiting_async_process(task)
1019-
if result is None:
1020-
self.scheduler.put_results(
1021-
[
1022-
RequestOutput(
1023-
request_id=task.request_id,
1024-
finished=True,
1025-
error_code=task.error_code,
1026-
error_msg=task.error_message,
1027-
)
1028-
]
1029-
)
1030-
need_check_req_ids.remove(task.request_id)
1031-
delete_tasks_list.append(task)
1032-
elif result is False:
1033-
if task.request_id in finished_ids:
1034-
need_check_req_ids.remove(task.request_id)
1035-
finished_ids.remove(task.request_id)
1036-
else:
1037-
time.sleep(0.001)
1038-
1039-
for tmp_task in delete_tasks_list:
1040-
tasks.remove(tmp_task)
1041-
# release resource in P
1042-
self.resource_manager.pre_recycle_resource(tmp_task.request_id)
1043-
1044-
# Fetch requests and add them to the scheduling queue
1045-
if tasks:
1046-
for task in tasks:
1047-
task.metrics.add_req_to_resource_manager_time = time.time()
1048-
trace_print(
1049-
LoggingEventName.RESOURCE_ALLOCATE_START, task.request_id, getattr(task, "user", "")
1050-
)
1051-
if self.cfg.scheduler_config.splitwise_role == "prefill":
1052-
self.resource_manager.add_request_in_p(tasks)
1053-
self.llm_logger.info(
1054-
f"P add requests into running queue: {[task.request_id for task in tasks]}"
1055-
)
1056-
else:
1057-
for task in tasks:
1058-
self.resource_manager.add_request(task)
1059-
is_fetching = False
1060-
except Exception as e:
1061-
self.llm_logger.error(f"fetching request error {e} {str(traceback.format_exc())}")
1062-
is_fetching = False
1063883

1064884
while self.running:
1065885
with self._pause_cond:
1066886
self._pause_cond.wait_for(lambda: not self.is_paused)
887+
1067888
try:
1068889
if self.engine_worker_queue.exist_tasks():
1069890
time.sleep(0.001)
1070891
continue
1071-
if self.cfg.scheduler_config.splitwise_role != "mixed":
1072-
if not is_fetching:
1073-
is_fetching = True
1074-
get_request_pool.submit(_fetch_request)
1075-
1076-
else:
1077-
if len(self.resource_manager.waiting) == 0 and (not is_fetching):
1078-
# Check if the thread pool is still available to avoid submitting tasks to a shutdown thread pool.
1079-
try:
1080-
is_fetching = True
1081-
get_request_pool.submit(_fetch_request)
1082-
except RuntimeError as e:
1083-
if "shutdown" in str(e):
1084-
self.llm_logger.info("Thread pool shutdown detected, exiting scheduler loop")
1085-
break
1086-
else:
1087-
raise
1088892

1089893
if hasattr(self.resource_manager, "scheduler_unhandled_request_num"):
1090894
self.resource_manager.scheduler_unhandled_request_num = self._get_scheduler_unhandled_request_num()
895+
1091896
# 2. Schedule requests
1092897
tasks, error_tasks = self.resource_manager.schedule()
1093898

@@ -2228,6 +2033,9 @@ def _exit_sub_services(self):
22282033
self.llm_logger.info("Exit sub services.....")
22292034
self.running = False
22302035

2036+
if hasattr(self, "_fetch_pool"):
2037+
self._fetch_pool.shutdown(wait=False)
2038+
22312039
if self.use_async_llm:
22322040
# Clean up worker processes first (before closing multiprocessing services)
22332041
if hasattr(self, "worker_proc") and self.worker_proc is not None:

0 commit comments

Comments
 (0)