From ada1e9d11180d2bbed464427f90d3267e05cb417 Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Thu, 24 Jul 2025 13:54:31 +0800 Subject: [PATCH 01/13] [LLM] support ep --- fastdeploy/engine/config.py | 4 ++ fastdeploy/engine/engine.py | 62 ++++++++++++++--------------- fastdeploy/engine/expert_service.py | 47 ++++++++++++---------- fastdeploy/worker/worker_process.py | 4 +- 4 files changed, 62 insertions(+), 55 deletions(-) diff --git a/fastdeploy/engine/config.py b/fastdeploy/engine/config.py index d8ebb38f000..cabe304c8fe 100644 --- a/fastdeploy/engine/config.py +++ b/fastdeploy/engine/config.py @@ -785,6 +785,10 @@ def postprocess(self): else: self.is_master = False + if self.tensor_parallel_size <= self.worker_num_per_node: + self.is_master = True + + import paddle self.paddle_commit_id = paddle.version.commit diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index 89070ae3e54..7ab75d2eb6f 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -243,38 +243,38 @@ def start(self, api_server_pid=None): self.splitwise_receive_thread.daemon = True self.splitwise_receive_thread.start() - self.cfg.init_cache_info() - - role = self.cfg.splitwise_role - host_ip = self.cfg.host_ip - disaggregate = self.cfg.disaggregate_info - if self.cfg.scheduler_config.name == "splitwise": - self.scheduler.start(role, host_ip, disaggregate) - - time.sleep(1) - - if self.cfg.parallel_config.enable_expert_parallel and self.cfg.parallel_config.data_parallel_size > 1: - self.dp_processed = [] - for i in range( - 1, - self.cfg.parallel_config.data_parallel_size // self.cfg.nnode, - ): - time.sleep(1) - self.dp_processed.append( - multiprocessing.Process( - target=start_expert_service, - args=( - self.cfg, - i + self.cfg.node_rank * self.cfg.worker_num_per_node, - self.ipc_signal_suffix, - ), - ) - ) - llm_logger.info( - f"Engine is initialized successfully with {self.cfg.tensor_parallel_size}" - + f" data parallel id {i}" + self.cfg.init_cache_info() + + role = self.cfg.splitwise_role + host_ip = self.cfg.host_ip + disaggregate = self.cfg.disaggregate_info + if self.cfg.scheduler_config.name == "splitwise": + self.scheduler.start(role, host_ip, disaggregate) + + time.sleep(1) + + if self.cfg.parallel_config.enable_expert_parallel and self.cfg.parallel_config.data_parallel_size > 1: + self.dp_processed = [] + for i in range( + 1, + self.cfg.parallel_config.data_parallel_size // self.cfg.nnode, + ): + time.sleep(1) + self.dp_processed.append( + multiprocessing.Process( + target=start_expert_service, + args=( + self.cfg, + i + self.cfg.node_rank * self.cfg.worker_num_per_node, + self.ipc_signal_suffix, + ), ) - self.dp_processed[-1].start() + ) + llm_logger.info( + f"Engine is initialized successfully with {self.cfg.tensor_parallel_size}" + + f" data parallel id {i}" + ) + self.dp_processed[-1].start() console_logger.info(f"Worker processes are launched with {time.time() - start_time} seconds.") return True diff --git a/fastdeploy/engine/expert_service.py b/fastdeploy/engine/expert_service.py index f2f5e9e17a8..313434e7387 100644 --- a/fastdeploy/engine/expert_service.py +++ b/fastdeploy/engine/expert_service.py @@ -50,10 +50,13 @@ def __init__(self, cfg, local_data_parallel_id): cfg (Config): Config object containing all the configuration parameters. """ self.cfg = cfg - start_pos = (local_data_parallel_id * self.cfg.tensor_parallel_size) % self.cfg.worker_num_per_node - end_pos = ((local_data_parallel_id + 1) * self.cfg.tensor_parallel_size) % self.cfg.worker_num_per_node - self.cfg.cache_config.rdma_comm_ports = self.cfg.cache_config.rdma_comm_ports[start_pos:end_pos] - self.cfg.local_device_ids = self.cfg.device_ids.split(",")[start_pos:end_pos] + start_pos = (local_data_parallel_id * self.cfg.tensor_parallel_size) % cfg.worker_num_per_node + end_pos = ((local_data_parallel_id + 1) * self.cfg.tensor_parallel_size) % cfg.worker_num_per_node + if cfg.splitwise_role != 'mixed': + self.cfg.cache_config.rdma_comm_ports = self.cfg.cache_config.rdma_comm_ports[ + start_pos:end_pos] + self.cfg.local_device_ids = self.cfg.device_ids.split( + ",")[start_pos:end_pos] self.cfg.parallel_config.local_data_parallel_id = local_data_parallel_id self.cfg.disaggregate_info = None @@ -78,11 +81,11 @@ def __init__(self, cfg, local_data_parallel_id): cfg.splitwise_role, local_data_parallel_id, ) - - if len(self.cfg.cache_config.pd_comm_port) == 1: - self.cfg.cache_config.pd_comm_port[0] = int(self.cfg.cache_config.pd_comm_port[0]) + local_data_parallel_id - else: - self.cfg.cache_config.pd_comm_port = [self.cfg.cache_config.pd_comm_port[local_data_parallel_id]] + if cfg.splitwise_role != 'mixed': + if len(self.cfg.cache_config.pd_comm_port) == 1: + self.cfg.cache_config.pd_comm_port[0] = int(self.cfg.cache_config.pd_comm_port[0]) + local_data_parallel_id + else: + self.cfg.cache_config.pd_comm_port = [self.cfg.cache_config.pd_comm_port[local_data_parallel_id]] self.split_connector = SplitwiseConnector( self.cfg, @@ -119,15 +122,16 @@ def start(self, ipc_signal_suffix, local_data_parallel_id): start_time = time.time() llm_logger.info(f"start expert service {local_data_parallel_id}") - - self.cache_manager_processes = self.resource_manager.cache_manager.launch_cache_manager( - cache_config=self.cfg.cache_config, - tensor_parallel_size=self.cfg.tensor_parallel_size, - device_ids=self.cfg.local_device_ids, - pod_ip=self.cfg.master_ip, - engine_worker_queue_port=self.cfg.engine_worker_queue_port, - pid_suffix=f"{local_data_parallel_id}_{ipc_signal_suffix}", - ) + if self.cfg.splitwise_role != 'mixed': + self.cache_manager_processes = self.resource_manager.cache_manager.launch_cache_manager( + cache_config=self.cfg.cache_config, + tensor_parallel_size=self.cfg.tensor_parallel_size, + device_ids=self.cfg.local_device_ids, + pod_ip=self.cfg.pod_ips[0], + engine_worker_queue_port=self.cfg.engine_worker_queue_port, + pid_suffix=f"{local_data_parallel_id}_{ipc_signal_suffix}" + ) + self.split_mode_get_tasks() self.insert_task_to_worker_thread = threading.Thread(target=self._insert_task_to_worker, args=()) self.insert_task_to_worker_thread.daemon = True @@ -138,7 +142,6 @@ def start(self, ipc_signal_suffix, local_data_parallel_id): self.token_processor.run() - self.split_mode_get_tasks() self.cfg.init_cache_info() @@ -321,13 +324,13 @@ def insert_tasks(self, tasks, current_id=-1, allocated=False): else: is_prefill = True self.token_processor.number_of_input_tokens += tasks[i].prompt_token_ids_len - - self.split_connector.send_cache_infos(tasks, current_id) + if is_decode or is_prefill: + self.split_connector.send_cache_infos(tasks, current_id) for task in tasks: task.infer_start_time = time.time() if not is_decode: llm_logger.info(f"Tasks are sent to engine, req_ids={req_ids}") - if not is_prefill: + if not is_prefill and self.cfg.cache_config.enable_chunked_prefill: if not self.cfg.enable_mm: self.update_requests_chunk_size(tasks) else: diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 8a12988c4f2..c9ec2194a8d 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -280,14 +280,14 @@ def event_loop_normal(self) -> None: paddle.distributed.barrier() self.insert_step = False - self.worker_healthy_live_signal.value[self.local_rank] = int(time.time()) + self.worker_healthy_live_signal.value[self.local_rank % self.max_chips_per_node] = int(time.time()) # The first worker detects whether there are tasks in the task queue if self.local_rank % mp_num_per_node == 0: if self.task_queue.num_tasks() > 0: # VL only support 1 batch to prefill if not self.fd_config.model_config.enable_mm or not self.worker.prefill_finished(): - if self.nnode > 1: + if self.nnode > 1 and self.parallel_config.tensor_parallel_size > 1 self.task_queue.read_finish_flag.set(1) else: self.exist_task_signal.value[self.fd_config.parallel_config.expert_parallel_rank] = 1 From 38b1efe181f95ec190c3f81eb4931c1751d6d7e6 Mon Sep 17 00:00:00 2001 From: ltd0924 <32387785+ltd0924@users.noreply.github.com> Date: Tue, 29 Jul 2025 11:41:37 +0800 Subject: [PATCH 02/13] Update worker_process.py --- fastdeploy/worker/worker_process.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index c9ec2194a8d..da52e88320f 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -287,7 +287,7 @@ def event_loop_normal(self) -> None: if self.task_queue.num_tasks() > 0: # VL only support 1 batch to prefill if not self.fd_config.model_config.enable_mm or not self.worker.prefill_finished(): - if self.nnode > 1 and self.parallel_config.tensor_parallel_size > 1 + if self.nnode > 1 and self.parallel_config.tensor_parallel_size > 1: self.task_queue.read_finish_flag.set(1) else: self.exist_task_signal.value[self.fd_config.parallel_config.expert_parallel_rank] = 1 From a9fa7f1fd70e203dfbcc6c2893e66325b6eff474 Mon Sep 17 00:00:00 2001 From: ltd0924 <32387785+ltd0924@users.noreply.github.com> Date: Tue, 29 Jul 2025 11:42:18 +0800 Subject: [PATCH 03/13] Update expert_service.py --- fastdeploy/engine/expert_service.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fastdeploy/engine/expert_service.py b/fastdeploy/engine/expert_service.py index 313434e7387..a7dbf179c04 100644 --- a/fastdeploy/engine/expert_service.py +++ b/fastdeploy/engine/expert_service.py @@ -51,7 +51,7 @@ def __init__(self, cfg, local_data_parallel_id): """ self.cfg = cfg start_pos = (local_data_parallel_id * self.cfg.tensor_parallel_size) % cfg.worker_num_per_node - end_pos = ((local_data_parallel_id + 1) * self.cfg.tensor_parallel_size) % cfg.worker_num_per_node + end_pos = start_pos + self.cfg.tensor_parallel_size if cfg.splitwise_role != 'mixed': self.cfg.cache_config.rdma_comm_ports = self.cfg.cache_config.rdma_comm_ports[ start_pos:end_pos] From ffa59fd4bd6bdc4884ea0685cbf6f2895ef2072c Mon Sep 17 00:00:00 2001 From: ltd0924 <32387785+ltd0924@users.noreply.github.com> Date: Tue, 29 Jul 2025 15:10:55 +0800 Subject: [PATCH 04/13] Update worker_process.py --- fastdeploy/worker/worker_process.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 68017118a13..0ff904fd42d 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -289,7 +289,7 @@ def event_loop_normal(self) -> None: if self.task_queue.num_tasks() > 0: # VL only support 1 batch to prefill - if not self.fd_config.model_config.enable_mm or not self.worker.prefill_finished(): + if not self.fd_config.model_config.enable_mm or not self.worker.exist_prefill(): if self.nnode > 1 and self.parallel_config.tensor_parallel_size > self.max_chips_per_node: self.task_queue.read_finish_flag.set(1) else: From ace1d9135c2452916dc22d7718ae62b261aabc67 Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Wed, 30 Jul 2025 16:18:37 +0800 Subject: [PATCH 05/13] format files --- fastdeploy/engine/config.py | 1 - fastdeploy/engine/expert_service.py | 19 +++++++++---------- 2 files changed, 9 insertions(+), 11 deletions(-) diff --git a/fastdeploy/engine/config.py b/fastdeploy/engine/config.py index 9fc587a60bc..63609cda27f 100644 --- a/fastdeploy/engine/config.py +++ b/fastdeploy/engine/config.py @@ -444,7 +444,6 @@ def postprocess(self): if self.tensor_parallel_size <= self.worker_num_per_node: self.is_master = True - import paddle self.paddle_commit_id = paddle.version.commit diff --git a/fastdeploy/engine/expert_service.py b/fastdeploy/engine/expert_service.py index a7dbf179c04..63b1b15beba 100644 --- a/fastdeploy/engine/expert_service.py +++ b/fastdeploy/engine/expert_service.py @@ -52,11 +52,9 @@ def __init__(self, cfg, local_data_parallel_id): self.cfg = cfg start_pos = (local_data_parallel_id * self.cfg.tensor_parallel_size) % cfg.worker_num_per_node end_pos = start_pos + self.cfg.tensor_parallel_size - if cfg.splitwise_role != 'mixed': - self.cfg.cache_config.rdma_comm_ports = self.cfg.cache_config.rdma_comm_ports[ - start_pos:end_pos] - self.cfg.local_device_ids = self.cfg.device_ids.split( - ",")[start_pos:end_pos] + if cfg.splitwise_role != "mixed": + self.cfg.cache_config.rdma_comm_ports = self.cfg.cache_config.rdma_comm_ports[start_pos:end_pos] + self.cfg.local_device_ids = self.cfg.device_ids.split(",")[start_pos:end_pos] self.cfg.parallel_config.local_data_parallel_id = local_data_parallel_id self.cfg.disaggregate_info = None @@ -81,9 +79,11 @@ def __init__(self, cfg, local_data_parallel_id): cfg.splitwise_role, local_data_parallel_id, ) - if cfg.splitwise_role != 'mixed': + if cfg.splitwise_role != "mixed": if len(self.cfg.cache_config.pd_comm_port) == 1: - self.cfg.cache_config.pd_comm_port[0] = int(self.cfg.cache_config.pd_comm_port[0]) + local_data_parallel_id + self.cfg.cache_config.pd_comm_port[0] = ( + int(self.cfg.cache_config.pd_comm_port[0]) + local_data_parallel_id + ) else: self.cfg.cache_config.pd_comm_port = [self.cfg.cache_config.pd_comm_port[local_data_parallel_id]] @@ -122,14 +122,14 @@ def start(self, ipc_signal_suffix, local_data_parallel_id): start_time = time.time() llm_logger.info(f"start expert service {local_data_parallel_id}") - if self.cfg.splitwise_role != 'mixed': + if self.cfg.splitwise_role != "mixed": self.cache_manager_processes = self.resource_manager.cache_manager.launch_cache_manager( cache_config=self.cfg.cache_config, tensor_parallel_size=self.cfg.tensor_parallel_size, device_ids=self.cfg.local_device_ids, pod_ip=self.cfg.pod_ips[0], engine_worker_queue_port=self.cfg.engine_worker_queue_port, - pid_suffix=f"{local_data_parallel_id}_{ipc_signal_suffix}" + pid_suffix=f"{local_data_parallel_id}_{ipc_signal_suffix}", ) self.split_mode_get_tasks() @@ -142,7 +142,6 @@ def start(self, ipc_signal_suffix, local_data_parallel_id): self.token_processor.run() - self.cfg.init_cache_info() role = self.cfg.splitwise_role From 9137785965e76854d7cdb3ae34f68bf641f60a6d Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Thu, 31 Jul 2025 13:26:30 +0800 Subject: [PATCH 06/13] optimize prefix cache --- fastdeploy/cache_manager/cache_messager.py | 179 ++++++++++++++++-- .../cache_manager/cache_transfer_manager.py | 97 ++-------- .../cache_manager/prefix_cache_manager.py | 110 +++++++++-- fastdeploy/engine/engine.py | 11 +- fastdeploy/spec_decode/mtp.py | 21 +- fastdeploy/worker/gpu_model_runner.py | 12 +- fastdeploy/worker/worker_process.py | 2 +- 7 files changed, 302 insertions(+), 130 deletions(-) diff --git a/fastdeploy/cache_manager/cache_messager.py b/fastdeploy/cache_manager/cache_messager.py index f11c406906d..1e3d4515708 100644 --- a/fastdeploy/cache_manager/cache_messager.py +++ b/fastdeploy/cache_manager/cache_messager.py @@ -17,16 +17,70 @@ import math import threading import time - +import argparse import numpy as np import paddle +import json +from fastdeploy.config import SpeculativeConfig from fastdeploy.cache_manager.transfer_factory import IPCCommManager, RDMACommManager from fastdeploy.inter_communicator import EngineWorkerQueue, IPCSignal from fastdeploy.utils import get_logger +from fastdeploy.model_executor.ops.gpu import set_data_ipc + + -logger = get_logger("cache_messager", "cache_messager.log") +def parse_args(): + """ + 从命令行解析参数 + """ + parser = argparse.ArgumentParser("Cache Messager") + parser.add_argument( + "--splitwise_role", + type=str, + default="mixed", + help="splitwise role, can be decode, prefill or mixed", + ) + parser.add_argument("--rank", type=int, default=0, help="current rank") + parser.add_argument("--device_id", type=int, default=0, help="device id") + parser.add_argument("--num_layers", type=int, default=1, help="model num layers") + parser.add_argument("--head_dim", type=int, default=1, help="model head dim") + parser.add_argument("--kv_num_head", type=int, default=1, help="model kv num head") + parser.add_argument("--rdma_port", type=str, default="", help="rmda port") + parser.add_argument("--mp_num", type=int, default=1, help="number of model parallel") + parser.add_argument("--engine_pid", type=str, default=None, help="engine pid") + parser.add_argument( + "--protocol", + type=str, + default="ipc", + help="cache transfer protocol, only surport ipc now", + ) + parser.add_argument("--pod_ip", type=str, default="0.0.0.0", help="pod ip") + parser.add_argument( + "--engine_worker_queue_port", + type=int, + default=9923, + help="engine worker queue port", + ) + parser.add_argument("--num_gpu_blocks", type=int, default=1, help="gpu cache block number") + parser.add_argument("--block_size", type=int, default=64, help="cache block size(tokens)") + parser.add_argument( + "--cache_dtype", + type=str, + default="bfloat16", + choices=["uint8", "bfloat16"], + help="cache dtype", + ) + parser.add_argument( + "--speculative_config", + type=json.loads, + default="{}", + help="speculative config", + ) + parser.add_argument("--local_data_parallel_id", type=int, default=0) + args = parser.parse_args() + return args class CacheMessager: """ @@ -83,7 +137,7 @@ def __init__( ) transfer_protocol = transfer_protocol.split(",") - logger.info(f"splitwise role: {splitwise_role}, {transfer_protocol}" f"rank: {rank}") + print(f"splitwise role: {splitwise_role}, {transfer_protocol}" f"rank: {rank}") # 1. initialize the cache_k_ptr_list and cache_v_ptr_list self.num_layers = num_layers @@ -108,7 +162,7 @@ def __init__( block_bytes = math.prod(cache_shape[1:]) if key_cache.dtype == paddle.bfloat16: block_bytes *= 2 - logger.info( + print( f"layers {num_layers} cache_shape: {cache_shape}, max_block_num: {max_block_num}, " f"block_bytes: {block_bytes}, dtype: {key_cache.dtype}" ) @@ -124,10 +178,10 @@ def __init__( cache_v, ) local_device_id = int(str(cache_k[0].place)[-2]) - logger.info(f"done create ipc_comm with local_device_id:{local_device_id}, ") + print(f"done create ipc_comm with local_device_id:{local_device_id}, ") elif protocol == "rdma": - logger.info(f"splitwise_role rdma: {self.splitwise_role}, rank: {self.rank}, gpu_id: {gpu_id}") + print(f"splitwise_role rdma: {self.splitwise_role}, rank: {self.rank}, gpu_id: {gpu_id}") self.messager[protocol] = RDMACommManager( splitwise_role, @@ -143,13 +197,9 @@ def __init__( self.gpu_id = gpu_id self.cache_info = dict() - layerwise_send_cache_thread = threading.Thread(target=self._prefill_layerwise_send_cache_thread) - layerwise_send_cache_thread.daemon = True - layerwise_send_cache_thread.start() - - logger.info(f"cache messager init finished, use {transfer_protocol}") + print(f"cache messager init finished, use {transfer_protocol}") - def _prefill_layerwise_send_cache_thread(self): + def prefill_layerwise_send_cache_thread(self): """ layerwise_send_cache_thread: send cache to other instance @@ -199,7 +249,7 @@ def _prefill_layerwise_send_cache_thread(self): cache_info = self.engine_worker_queue.get_cache_info() if cache_info: - logger.debug(f"cache info {cache_info}") + print(f"cache info {cache_info}") for info in cache_info: if info["request_id"] in self.cache_info: self.cache_info[info["request_id"]].update(info) @@ -211,7 +261,7 @@ def _prefill_layerwise_send_cache_thread(self): current_info["src_block_ids"] = current_src_blocks current_info["current_layer_ids"] = 0 current_info["status"] = "init" - logger.info(f"start cache_infos: {current_info}") + print(f"start cache_infos: {current_info}") self.cache_info[info["request_id"]] = current_info self.last_step_idx = min(self.last_step_idx, current_info["current_id"]) else: @@ -229,7 +279,7 @@ def _prefill_layerwise_send_cache_thread(self): if not self.cache_info: time.sleep(0.001) continue - logger.debug(f"prefilled_layer_idx: {prefilled_layer_idx}, prefilled_step_idx: {prefilled_step_idx}") + print(f"prefilled_layer_idx: {prefilled_layer_idx}, prefilled_step_idx: {prefilled_step_idx}") for req_id, item in list(self.cache_info.items()): if "status" not in item: continue @@ -246,7 +296,7 @@ def _prefill_layerwise_send_cache_thread(self): target_id = int(item["rdma_ports"][self.rank]) status = self.messager[current_transfer_protocol].connect(target_ip, target_id) if not status: - logger.error(f"connect to {target_ip}:{target_id} failed") + print(f"connect to {target_ip}:{target_id} failed") item["status"] = "error" self.engine_worker_queue.finish_request_barrier.wait() if self.rank == 0: @@ -276,7 +326,7 @@ def _prefill_layerwise_send_cache_thread(self): self.engine_worker_queue.finish_request_barrier.wait() if self.rank == 0: self.engine_worker_queue.put_finished_req([(item["request_id"], "write cache error")]) - logger.error( + print( f"write cache failed, layer_idx: {layer_idx}, " f"req_id: {item['request_id']}, dest_ip: {target_ip}" ) @@ -287,7 +337,7 @@ def _prefill_layerwise_send_cache_thread(self): block_num = len(src_block_ids) avg_time_per_block = cost_time * 1000 / block_num # ms send_cache_speed = block_num * self.block_bytes / 1073741824 / cost_time # GB/s - logger.debug( + print( f"finish write cache for a layer, {item['request_id']}, {layer_idx}" f" {current_transfer_protocol}" f"block_num: {block_num}, send_cache_speed(GB/s): {round(send_cache_speed, 5)}," @@ -297,15 +347,102 @@ def _prefill_layerwise_send_cache_thread(self): if item["layer_idx"] == self.num_layers: if item["transfer_protocol"] == "ipc": self.messager["ipc"].write_block_by_sync(target_id) - logger.info(f"finish write cache {item['request_id']}") + print(f"finish write cache {item['request_id']}") self.engine_worker_queue.finish_request_barrier.wait() if self.rank == 0: self.engine_worker_queue.put_finished_req([(item["request_id"], "finished")]) - logger.info(f"put write cache {item['request_id']}") + print(f"put write cache {item['request_id']}") del self.cache_info[req_id] self.last_step_idx = prefilled_step_idx self.last_layer_idx = prefilled_layer_idx except Exception as e: - logger.error(f"prefill layerwise send cache thread has exception: {e}") + print(f"prefill layerwise send cache thread has exception: {e}") + + +def main(): + device = args.device_id + rank = args.rank + paddle.set_device(f"gpu:{device}") + cache_type = args.cache_dtype + speculative_config = SpeculativeConfig(args.speculative_config) + num_extra_layers = speculative_config.num_extra_cache_layer + num_extra_layer_gpu_blocks = int(args.num_gpu_blocks * speculative_config.num_gpu_block_expand_ratio) + gpu_cache_kvs = {} + gpu_cache_k_tensors = [] + gpu_cache_v_tensors = [] + + for i in range(args.num_layers + num_extra_layers): + num_gpu_blocks = args.num_gpu_blocks if i < args.num_layers else num_extra_layer_gpu_blocks + + gpu_cache_kvs[f"key_caches_{i}_rank{rank}_device{device}"] = paddle.full( + shape=[ + num_gpu_blocks, + args.kv_num_head, + args.block_size, + args.head_dim, + ], + fill_value=0, + dtype=cache_type, + ) + gpu_cache_k_tensors.append(gpu_cache_kvs[f"key_caches_{i}_rank{rank}_device{device}"]) + gpu_cache_kvs[f"value_caches_{i}_rank{rank}_device{device}"] = paddle.full( + shape=[ + num_gpu_blocks, + args.kv_num_head, + args.block_size, + args.head_dim, + ], + fill_value=0, + dtype=cache_type, + ) + gpu_cache_v_tensors.append(gpu_cache_kvs[f"value_caches_{i}_rank{rank}_device{device}"]) + + set_data_ipc( + gpu_cache_kvs[f"key_caches_{i}_rank{rank}_device{device}"], + f"key_caches_{i}_rank{rank}.device{device}", + ) + set_data_ipc( + gpu_cache_kvs[f"value_caches_{i}_rank{rank}_device{device}"], + f"value_caches_{i}_rank{rank}.device{device}", + ) + cache_kv_size_byte = sum([tmp.numel() * 1 for key, tmp in gpu_cache_kvs.items()]) + print(f"device :{device}") + print(f"cache_kv_size_byte : {cache_kv_size_byte}") + print(f"done init cache (full) gmem alloc : {paddle.device.cuda.memory_allocated()}") + + cache_messager = CacheMessager( + splitwise_role=args.splitwise_role, + transfer_protocol=args.protocol, + pod_ip=args.pod_ip, + engine_worker_queue_port=args.engine_worker_queue_port, + local_data_parallel_id=args.local_data_parallel_id, + gpu_cache_kvs=gpu_cache_kvs, + rank=rank, + nranks=args.mp_num, + num_layers=args.num_layers + num_extra_layers, + gpu_id=device, + rdma_port=args.rdma_port, + ) + + cache_ready_signal_data = np.zeros(shape=[args.mp_num], dtype=np.int32) + cache_ready_signal = IPCSignal( + name="cache_ready_signal", + array=cache_ready_signal_data, + dtype=np.int32, + suffix=args.engine_pid, + create=False, + ) + cache_ready_signal.value[rank] = 1 + cache_messager.prefill_layerwise_send_cache_thread() + + +if __name__ == "__main__": + + args = parse_args() + + print("create cache messager...") + print(f"{args}") + main() + diff --git a/fastdeploy/cache_manager/cache_transfer_manager.py b/fastdeploy/cache_manager/cache_transfer_manager.py index 34ccf144ca8..7b8e576cc20 100644 --- a/fastdeploy/cache_manager/cache_transfer_manager.py +++ b/fastdeploy/cache_manager/cache_transfer_manager.py @@ -28,8 +28,9 @@ from fastdeploy.inter_communicator import EngineCacheQueue, IPCSignal from fastdeploy.model_executor.ops.gpu import ( cuda_host_alloc, - set_data_ipc, + share_external_data, swap_cache_all_layers, + ) from fastdeploy.utils import get_logger @@ -39,26 +40,12 @@ def parse_args(): 从命令行解析参数 """ parser = argparse.ArgumentParser("Cache transfer manager") - parser.add_argument( - "--splitwise_role", - type=str, - default="mixed", - help="splitwise role, can be decode, prefill or mixed", - ) parser.add_argument("--rank", type=int, default=0, help="current rank") parser.add_argument("--device_id", type=int, default=0, help="device id") parser.add_argument("--num_layers", type=int, default=1, help="model num layers") parser.add_argument("--head_dim", type=int, default=1, help="model head dim") parser.add_argument("--kv_num_head", type=int, default=1, help="model kv num head") - parser.add_argument("--rdma_port", type=str, default="", help="rmda port") parser.add_argument("--mp_num", type=int, default=1, help="number of model parallel") - parser.add_argument( - "--protocol", - type=str, - default="ipc", - help="cache transfer protocol, only surport ipc now", - ) - parser.add_argument("--enable_splitwise", type=int, default=0, help="enable splitwise ") parser.add_argument("--cache_queue_port", type=int, default=9923, help="cache queue port") parser.add_argument("--pod_ip", type=str, default="0.0.0.0", help="pod ip") parser.add_argument( @@ -68,7 +55,6 @@ def parse_args(): help="engine worker queue port", ) parser.add_argument("--engine_pid", type=str, default=None, help="engine pid") - parser.add_argument("--num_gpu_blocks", type=int, default=1, help="gpu cache block number") parser.add_argument("--num_cpu_blocks", type=int, default=4, help="cpu cache block number") parser.add_argument("--block_size", type=int, default=64, help="cache block size(tokens)") @@ -109,7 +95,6 @@ def __init__(self, args): device = args.device_id rank = args.rank - paddle.set_device(f"gpu:{device}") self.gpu_cache_kvs = {} self.cpu_cache_kvs = {} self.gpu_cache_k_tensors = [] @@ -138,40 +123,26 @@ def __init__(self, args): self.num_cpu_blocks = args.num_cpu_blocks cache_type = args.cache_dtype - for i in range(args.num_layers + self.num_extra_layers): - num_gpu_blocks = args.num_gpu_blocks if i < args.num_layers else self.num_extra_layer_gpu_blocks - - self.gpu_cache_kvs[f"key_caches_{i}_rank{rank}_device{device}"] = paddle.full( - shape=[ - num_gpu_blocks, + cache_shape = [ + args.num_gpu_blocks, args.kv_num_head, args.block_size, args.head_dim, - ], - fill_value=0, - dtype=cache_type, - ) - self.gpu_cache_k_tensors.append(self.gpu_cache_kvs[f"key_caches_{i}_rank{rank}_device{device}"]) - self.gpu_cache_kvs[f"value_caches_{i}_rank{rank}_device{device}"] = paddle.full( - shape=[ - num_gpu_blocks, - args.kv_num_head, - args.block_size, - args.head_dim, - ], - fill_value=0, - dtype=cache_type, - ) - self.gpu_cache_v_tensors.append(self.gpu_cache_kvs[f"value_caches_{i}_rank{rank}_device{device}"]) + ] + + for i in range(args.num_layers + self.num_extra_layers): + num_gpu_blocks = args.num_gpu_blocks if i < args.num_layers else self.num_extra_layer_gpu_blocks + key_name = f"key_caches_{i}_rank{rank}.device{device}" + value_name = f"value_caches_{i}_rank{rank}.device{device}" + key_cache = paddle.empty(shape=[], dtype=cache_type) + value_cache = paddle.empty(shape=[], dtype=cache_type) + key_cache = share_external_data(key_cache, key_name, cache_shape) + value_cache = share_external_data(value_cache, value_name, cache_shape) + self.gpu_cache_kvs[key_name] = key_cache + self.gpu_cache_kvs[value_name] = value_cache + self.gpu_cache_k_tensors.append(self.gpu_cache_kvs[key_name]) + self.gpu_cache_v_tensors.append(self.gpu_cache_kvs[value_name]) - set_data_ipc( - self.gpu_cache_kvs[f"key_caches_{i}_rank{rank}_device{device}"], - f"key_caches_{i}_rank{rank}.device{device}", - ) - set_data_ipc( - self.gpu_cache_kvs[f"value_caches_{i}_rank{rank}_device{device}"], - f"value_caches_{i}_rank{rank}.device{device}", - ) cache_kv_size_byte = sum([tmp.numel() * 1 for key, tmp in self.gpu_cache_kvs.items()]) logger.info(f"device :{self.device}") logger.info(f"cache_kv_size_byte : {cache_kv_size_byte}") @@ -190,37 +161,7 @@ def __init__(self, args): ) self.v_dst_ptrs.append(self.cpu_cache_kvs[f"value_caches_{i}_rank{rank}"]) - cache_ready_signal_data = np.zeros(shape=[args.mp_num], dtype=np.int32) - self.cache_ready_signal = IPCSignal( - name="cache_ready_signal", - array=cache_ready_signal_data, - dtype=np.int32, - suffix=args.engine_pid, - create=False, - ) - self.cache_ready_signal.value[self.rank] = 1 - - paddle.set_device(f"gpu:{device}") - if args.enable_splitwise: - logger.debug("create cache messager...") - logger.info(f"{args}") - from fastdeploy.cache_manager.cache_messager import CacheMessager - - self.cache_messager = CacheMessager( - splitwise_role=args.splitwise_role, - transfer_protocol=args.protocol, - pod_ip=args.pod_ip, - engine_worker_queue_port=args.engine_worker_queue_port, - local_data_parallel_id=args.local_data_parallel_id, - gpu_cache_kvs=self.gpu_cache_kvs, - rank=self.rank, - nranks=args.mp_num, - num_layers=args.num_layers + self.num_extra_layers, - gpu_id=self.device, - rdma_port=args.rdma_port, - ) - logger.info("successfully create cache messager") - logger.info(f"done init CacheMessager gmem alloc : {paddle.device.cuda.memory_allocated()}") + cache_task_broadcast_data = np.zeros(shape=[1], dtype=np.int32) self.cache_task_broadcast_signal = IPCSignal( diff --git a/fastdeploy/cache_manager/prefix_cache_manager.py b/fastdeploy/cache_manager/prefix_cache_manager.py index dd191c87f08..8f258f4989b 100644 --- a/fastdeploy/cache_manager/prefix_cache_manager.py +++ b/fastdeploy/cache_manager/prefix_cache_manager.py @@ -140,6 +140,84 @@ def launch_cache_manager( filename = "cache_transfer_manager.py" py_path = os.path.join(current_dir_path, filename) + + cache_messager_processes = [] + if self.splitwise_role != "mixed": + cache_messager_processes = self.launch_cache_messager( + cache_config, + tensor_parallel_size, + device_ids, + pod_ip, + engine_worker_queue_port, + pid_suffix, + ) + if cache_messager_processes is None: + raise RuntimeError("Launch cache messager failed") + return [] + + if ( + hasattr(cache_config.model_cfg, "num_key_value_heads") + and hasattr(cache_config.model_cfg, "num_key_value_heads") + and cache_config.model_cfg.num_key_value_heads is not None + and int(cache_config.model_cfg.num_key_value_heads) > 0 + ): + kv_num_head = int(cache_config.model_cfg.num_key_value_heads) // tensor_parallel_size + else: + kv_num_head = cache_config.model_cfg.num_attention_heads // tensor_parallel_size + + + log_dir = envs.FD_LOG_DIR + cache_manager_processes = [] + for i in range(tensor_parallel_size): + launch_cmd = (f" {sys.executable} {py_path}" + + f" --device_id {int(device_ids[i])}" + + f" --rank {i}" + + f" --num_layers {cache_config.model_cfg.num_layers}" + + f" --head_dim {cache_config.model_cfg.head_dim}" + + f" --kv_num_head {kv_num_head}" + + f" --mp_num {tensor_parallel_size}" + + f" --cache_dtype {cache_config.cache_dtype}" + + f" --cache_queue_port {cache_config.cache_queue_port}" + + f" --pod_ip {pod_ip}" + + f" --engine_worker_queue_port {engine_worker_queue_port}" + + f" --num_gpu_blocks {cache_config.total_block_num}" + + f" --num_cpu_blocks {cache_config.num_cpu_blocks}" + + f" --bytes_per_layer_per_block {cache_config.bytes_per_layer_per_block}" + + f" --block_size {cache_config.block_size}" + + f" --engine_pid {pid_suffix}" + + f" --local_data_parallel_id {self.local_data_parallel_id}" + + f" --speculative_config '{self.speculative_config.to_json_string()}'" + + f" >{log_dir}/launch_cache_manager_{int(device_ids[i])}.log 2>&1" + ) + logger.info(f"Launch cache transfer manager, command:{launch_cmd}") + cache_manager_processes.append(subprocess.Popen(launch_cmd, shell=True, preexec_fn=os.setsid)) + exit_code = cache_manager_processes[-1].poll() + if exit_code is None: + logger.info("Launch cache transfer manager successful") + else: + logger.info("Launch cache transfer manager failed, see launch_cache_manager.log for more information") + + if cache_config.enable_hierarchical_cache and self.num_cpu_blocks > 0: + logger.info("Enable hierarchical cache.") + self._enable_cpu_cache() + cache_manager_processes.extend(cache_messager_processes) + return cache_manager_processes + + + + def launch_cache_messager(self, + cache_config, + tensor_parallel_size, + device_ids, + pod_ip, + engine_worker_queue_port, + pid_suffix + ): + """ + launch_cache_messager function used to initialize the cache messager. + """ + current_dir_path = os.path.split(os.path.abspath(__file__))[0] + filename = "cache_messager.py" if ( hasattr(cache_config.model_cfg, "num_key_value_heads") and hasattr(cache_config.model_cfg, "num_key_value_heads") @@ -158,8 +236,10 @@ def launch_cache_manager( suffix=pid_suffix, create=True, ) + + py_path = os.path.join(current_dir_path, filename) log_dir = envs.FD_LOG_DIR - cache_manager_processes = [] + cache_messager_processes = [] for i in range(tensor_parallel_size): launch_cmd = ( "FLAGS_allocator_strategy=auto_growth CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7" @@ -173,37 +253,31 @@ def launch_cache_manager( + f" --kv_num_head {kv_num_head}" + f" --mp_num {tensor_parallel_size}" + f" --cache_dtype {cache_config.cache_dtype}" - + f" --cache_queue_port {cache_config.cache_queue_port}" - + f" --enable_splitwise {int(self.enable_splitwise)}" + f" --pod_ip {pod_ip}" + f" --engine_worker_queue_port {engine_worker_queue_port}" + f" --num_gpu_blocks {cache_config.total_block_num}" - + f" --num_cpu_blocks {cache_config.num_cpu_blocks}" - + f" --bytes_per_layer_per_block {cache_config.bytes_per_layer_per_block}" + f" --block_size {cache_config.block_size}" - + f" --engine_pid {pid_suffix}" + f" --protocol {cache_config.cache_transfer_protocol}" + f" --local_data_parallel_id {self.local_data_parallel_id}" + + f" --engine_pid {pid_suffix}" + f" --rdma_port {cache_config.rdma_comm_ports[i] if cache_config.rdma_comm_ports is not None else '0'}" + f" --speculative_config '{self.speculative_config.to_json_string()}'" - + f" >{log_dir}/launch_cache_manager_{int(device_ids[i])}.log 2>&1" + + f" >{log_dir}/launch_cache_messager_{int(device_ids[i])}.log 2>&1" ) - logger.info(f"Launch cache transfer manager, command:{launch_cmd}") - cache_manager_processes.append(subprocess.Popen(launch_cmd, shell=True, preexec_fn=os.setsid)) - # 等待cache初始化完毕 - logger.info("Waiting for cache transfer manager ready...") + logger.info(f"Launch cache messager, command:{launch_cmd}") + cache_messager_processes.append(subprocess.Popen(launch_cmd, shell=True, preexec_fn=os.setsid)) + logger.info("Waiting for cache ready...") while np.sum(self.cache_ready_signal.value) != tensor_parallel_size: time.sleep(1) - exit_code = cache_manager_processes[-1].poll() + exit_code = cache_messager_processes[-1].poll() if exit_code is None: - logger.info("Launch cache transfer manager successful") + logger.info("Launch cache messager successful") else: - logger.info("Launch cache transfer manager failed, see launch_cache_manager.log for more information") + logger.info("Launch cache messager failed, see launch_cache_messager.log for more information") + cache_messager_processes = None + return cache_messager_processes + - if cache_config.enable_hierarchical_cache and self.num_cpu_blocks > 0: - logger.info("Enable hierarchical cache.") - self._enable_cpu_cache() - return cache_manager_processes def update_cache_config(self, cache_config): """ diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index 9ddf0cbf775..771c833936a 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -816,12 +816,11 @@ def insert_tasks(self, tasks, current_id=-1, allocated=False): if not is_decode: llm_logger.info(f"Tasks are sent to engine, req_ids={req_ids}") for task in tasks: - task.inference_start_time = time.time() - if not is_prefill: - if not self.cfg.enable_mm: - self.update_requests_chunk_size(tasks) - else: - self.update_mm_requests_chunk_size(tasks) + task.inference_start_time = time.time(): + if not self.cfg.enable_mm: + self.update_requests_chunk_size(tasks) + else: + self.update_mm_requests_chunk_size(tasks) self.engine_worker_queue.put_tasks((tasks, self.resource_manager.real_bsz)) if is_prefill and self.cfg.scheduler_config.name != "splitwise": self.engine_worker_queue.available_prefill_instances.put(1) diff --git a/fastdeploy/spec_decode/mtp.py b/fastdeploy/spec_decode/mtp.py index 39f0fce4272..c9d0d52d849 100644 --- a/fastdeploy/spec_decode/mtp.py +++ b/fastdeploy/spec_decode/mtp.py @@ -38,6 +38,7 @@ mtp_save_first_token, mtp_step_paddle, share_external_data, + set_data_ipc ) from fastdeploy.model_executor.pre_and_post_process import pre_process, rebuild_padding @@ -141,9 +142,7 @@ def initialize_kv_cache(self): kv_cache_shape = self.attn_backends[0].get_kv_cache_shape( max_num_blocks=self.num_gpu_blocks, kv_cache_quant_type=kv_cache_quant_type ) - if not self.parallel_config.do_profile and ( - self.cache_config.enable_prefix_caching or self.parallel_config.splitwise_role != "mixed" - ): + if not self.parallel_config.do_profile and self.parallel_config.splitwise_role != "mixed": cache_kvs_list = [] for i in range( self.num_main_model_layers, @@ -160,7 +159,10 @@ def initialize_kv_cache(self): self.model_inputs["caches"] = cache_kvs_list else: - for i in range(self.model_config.num_hidden_layers): + for i in range( + self.num_main_model_layers, + self.num_main_model_layers + self.model_config.num_hidden_layers, + ): self.cache_kvs[f"key_caches_{i}"] = paddle.full( shape=kv_cache_shape, fill_value=0, @@ -171,6 +173,15 @@ def initialize_kv_cache(self): fill_value=0, dtype=cache_type, ) + if self.cache_config.enable_prefix_caching: + set_data_ipc( + self.cache_kvs[f"key_caches_{i}"], + f"key_caches_{i}_rank{self.local_rank}.device{self.device_id}", + ) + set_data_ipc( + self.cache_kvs[f"value_caches_{i}"], + f"value_caches_{i}_rank{self.local_rank}.device{self.device_id}", + ) self.model_inputs["caches"] = list(self.cache_kvs.values()) for value in self.cache_kvs.values(): del value @@ -235,7 +246,7 @@ def update_block_num(self, num_gpu_blocks) -> None: self.main_model_num_gpu_blocks = num_gpu_blocks self.num_gpu_blocks = int(num_gpu_blocks * self.speculative_config.num_gpu_block_expand_ratio) - if not (self.cache_config.enable_prefix_caching or self.parallel_config.splitwise_role != "mixed"): + if self.parallel_config.splitwise_role == "mixed": self.initialize_kv_cache() # Reset free list diff --git a/fastdeploy/worker/gpu_model_runner.py b/fastdeploy/worker/gpu_model_runner.py index 4b67b595e84..c037926634e 100644 --- a/fastdeploy/worker/gpu_model_runner.py +++ b/fastdeploy/worker/gpu_model_runner.py @@ -45,6 +45,7 @@ recover_decode_task, set_value_by_flags_and_idx, share_external_data, + set_data_ipc ) from fastdeploy.model_executor.pre_and_post_process import ( post_process, @@ -904,7 +905,7 @@ def initialize_kv_cache(self, profile: bool = False) -> None: ) local_rank = self.local_rank % self.parallel_config.tensor_parallel_size - if not profile and (self.cache_config.enable_prefix_caching or self.parallel_config.splitwise_role != "mixed"): + if not profile and self.parallel_config.splitwise_role != "mixed": cache_kvs_list = [] for i in range(self.model_config.num_hidden_layers): key_cache = paddle.empty(shape=[], dtype=cache_type) @@ -930,6 +931,15 @@ def initialize_kv_cache(self, profile: bool = False) -> None: fill_value=0, dtype=cache_type, ) + if self.cache_config.enable_prefix_caching: + set_data_ipc( + cache_kvs[f"key_caches_{i}"], + f"key_caches_{i}_rank{local_rank}.device{self.device_id}", + ) + set_data_ipc( + cache_kvs[f"value_caches_{i}"], + f"value_caches_{i}_rank{local_rank}.device{self.device_id}", + ) self.share_inputs["caches"] = list(cache_kvs.values()) for value in cache_kvs.values(): del value diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 54f7019c871..c22e1d979b3 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -408,7 +408,7 @@ def initialize_kv_cache(self) -> None: logger.info(f"------- num_blocks_global: {num_blocks_local} --------") # wait engine launch cache_manager - if self.cache_config.enable_prefix_caching or self.parallel_config.splitwise_role != "mixed": + if self.parallel_config.splitwise_role != "mixed": launched_cache_manager_signal_data = np.zeros([1], dtype=np.int32) self.launched_cache_manager_signal = IPCSignal( name="launched_cache_manager_signal", From b02d7c3feaebd7c028718510726674197bc6bceb Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Thu, 31 Jul 2025 13:38:41 +0800 Subject: [PATCH 07/13] optimize prefix cache --- fastdeploy/engine/engine.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index 771c833936a..d5982c687ce 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -816,7 +816,7 @@ def insert_tasks(self, tasks, current_id=-1, allocated=False): if not is_decode: llm_logger.info(f"Tasks are sent to engine, req_ids={req_ids}") for task in tasks: - task.inference_start_time = time.time(): + task.inference_start_time = time.time() if not self.cfg.enable_mm: self.update_requests_chunk_size(tasks) else: From aecd08c06552ade423af205ab3c4f7e47fbb62b8 Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Thu, 31 Jul 2025 20:21:26 +0800 Subject: [PATCH 08/13] optimize prefix cache --- fastdeploy/cache_manager/cache_messager.py | 24 +++++++++---------- .../cache_manager/cache_transfer_manager.py | 9 +++---- .../cache_manager/prefix_cache_manager.py | 4 ++-- fastdeploy/engine/engine.py | 10 ++++---- 4 files changed, 25 insertions(+), 22 deletions(-) diff --git a/fastdeploy/cache_manager/cache_messager.py b/fastdeploy/cache_manager/cache_messager.py index 1e3d4515708..11853823ca1 100644 --- a/fastdeploy/cache_manager/cache_messager.py +++ b/fastdeploy/cache_manager/cache_messager.py @@ -43,7 +43,7 @@ def parse_args(): ) parser.add_argument("--rank", type=int, default=0, help="current rank") parser.add_argument("--device_id", type=int, default=0, help="device id") - parser.add_argument("--num_layers", type=int, default=1, help="model num layers") + parser.add_argument("--num_hidden_layers", type=int, default=1, help="model num layers") parser.add_argument("--head_dim", type=int, default=1, help="model head dim") parser.add_argument("--kv_num_head", type=int, default=1, help="model kv num head") parser.add_argument("--rdma_port", type=str, default="", help="rmda port") @@ -97,7 +97,7 @@ def __init__( gpu_cache_kvs, rank, nranks, - num_layers, + num_hidden_layers, gpu_id=0, rdma_port=None, ): @@ -111,7 +111,7 @@ def __init__( gpu_cache_kvs (dict): GPU kv cache rank (int): current rank nranks (int): global rank number - num_layers (int): model layer number + num_hidden_layers (int): model layer number gpu_id (int, optional): GPU ID rdma_port (int, optional): RDMA port @@ -140,13 +140,13 @@ def __init__( print(f"splitwise role: {splitwise_role}, {transfer_protocol}" f"rank: {rank}") # 1. initialize the cache_k_ptr_list and cache_v_ptr_list - self.num_layers = num_layers + self.num_hidden_layers = num_hidden_layers cache_k_ptr_list = [] cache_v_ptr_list = [] cache_k = [] cache_v = [] self.messager = {} - for layer_idx in range(self.num_layers): + for layer_idx in range(self.num_hidden_layers): key_cache = self.gpu_cache_kvs[f"key_caches_{layer_idx}_rank{self.rank}_device{gpu_id}"] val_cache = self.gpu_cache_kvs[f"value_caches_{layer_idx}_rank{self.rank}_device{gpu_id}"] cache_k.append(key_cache) @@ -163,7 +163,7 @@ def __init__( if key_cache.dtype == paddle.bfloat16: block_bytes *= 2 print( - f"layers {num_layers} cache_shape: {cache_shape}, max_block_num: {max_block_num}, " + f"layers {num_hidden_layers} cache_shape: {cache_shape}, max_block_num: {max_block_num}, " f"block_bytes: {block_bytes}, dtype: {key_cache.dtype}" ) self.block_bytes = block_bytes @@ -268,7 +268,7 @@ def prefill_layerwise_send_cache_thread(self): self.cache_info[info["request_id"]] = info prefilled_layer_idx = layer_shm_value.value[0] prefilled_step_idx = step_shm_value.value[0] - if prefilled_layer_idx == self.num_layers - 1: + if prefilled_layer_idx == self.num_hidden_layers - 1: time.sleep(0.001) prefilled_layer_idx = layer_shm_value.value[0] prefilled_step_idx = step_shm_value.value[0] @@ -308,7 +308,7 @@ def prefill_layerwise_send_cache_thread(self): src_block_ids = paddle.to_tensor(item["src_block_ids"], dtype="int32", place="cpu") dest_block_ids = paddle.to_tensor(item["dest_block_ids"], dtype="int32", place="cpu") if item["current_id"] < prefilled_step_idx: - current_layer_idx = self.num_layers + current_layer_idx = self.num_hidden_layers else: current_layer_idx = prefilled_layer_idx + 1 @@ -344,7 +344,7 @@ def prefill_layerwise_send_cache_thread(self): f"avg_time per block(ms): {round(avg_time_per_block, 5)}" ) item["layer_idx"] = current_layer_idx - if item["layer_idx"] == self.num_layers: + if item["layer_idx"] == self.num_hidden_layers: if item["transfer_protocol"] == "ipc": self.messager["ipc"].write_block_by_sync(target_id) print(f"finish write cache {item['request_id']}") @@ -373,8 +373,8 @@ def main(): gpu_cache_k_tensors = [] gpu_cache_v_tensors = [] - for i in range(args.num_layers + num_extra_layers): - num_gpu_blocks = args.num_gpu_blocks if i < args.num_layers else num_extra_layer_gpu_blocks + for i in range(args.num_hidden_layers + num_extra_layers): + num_gpu_blocks = args.num_gpu_blocks if i < args.num_hidden_layers else num_extra_layer_gpu_blocks gpu_cache_kvs[f"key_caches_{i}_rank{rank}_device{device}"] = paddle.full( shape=[ @@ -421,7 +421,7 @@ def main(): gpu_cache_kvs=gpu_cache_kvs, rank=rank, nranks=args.mp_num, - num_layers=args.num_layers + num_extra_layers, + num_hidden_layers=args.num_hidden_layers + num_extra_layers, gpu_id=device, rdma_port=args.rdma_port, ) diff --git a/fastdeploy/cache_manager/cache_transfer_manager.py b/fastdeploy/cache_manager/cache_transfer_manager.py index 7b8e576cc20..2e2d8b1bde8 100644 --- a/fastdeploy/cache_manager/cache_transfer_manager.py +++ b/fastdeploy/cache_manager/cache_transfer_manager.py @@ -42,7 +42,7 @@ def parse_args(): parser = argparse.ArgumentParser("Cache transfer manager") parser.add_argument("--rank", type=int, default=0, help="current rank") parser.add_argument("--device_id", type=int, default=0, help="device id") - parser.add_argument("--num_layers", type=int, default=1, help="model num layers") + parser.add_argument("--num_hidden_layers", type=int, default=1, help="model num layers") parser.add_argument("--head_dim", type=int, default=1, help="model head dim") parser.add_argument("--kv_num_head", type=int, default=1, help="model kv num head") parser.add_argument("--mp_num", type=int, default=1, help="number of model parallel") @@ -130,8 +130,9 @@ def __init__(self, args): args.head_dim, ] - for i in range(args.num_layers + self.num_extra_layers): - num_gpu_blocks = args.num_gpu_blocks if i < args.num_layers else self.num_extra_layer_gpu_blocks + for i in range(args.num_hidden_layers + self.num_extra_layers): + num_gpu_blocks = args.num_gpu_blocks if i < args.num_hidden_layers else self.num_extra_layer_gpu_blocks + cache_shape[0] = num_gpu_blocks key_name = f"key_caches_{i}_rank{rank}.device{device}" value_name = f"value_caches_{i}_rank{rank}.device{device}" key_cache = paddle.empty(shape=[], dtype=cache_type) @@ -151,7 +152,7 @@ def __init__(self, args): paddle.set_device("cpu") self.k_dst_ptrs = [] self.v_dst_ptrs = [] - for i in range(args.num_layers + self.num_extra_layers): + for i in range(args.num_hidden_layers + self.num_extra_layers): self.cpu_cache_kvs[f"key_caches_{i}_rank{rank}"] = cuda_host_alloc( args.num_cpu_blocks * args.bytes_per_layer_per_block ) diff --git a/fastdeploy/cache_manager/prefix_cache_manager.py b/fastdeploy/cache_manager/prefix_cache_manager.py index 8f258f4989b..6bbf5511b39 100644 --- a/fastdeploy/cache_manager/prefix_cache_manager.py +++ b/fastdeploy/cache_manager/prefix_cache_manager.py @@ -172,7 +172,7 @@ def launch_cache_manager( launch_cmd = (f" {sys.executable} {py_path}" + f" --device_id {int(device_ids[i])}" + f" --rank {i}" - + f" --num_layers {cache_config.model_cfg.num_layers}" + + f" --num_hidden_layers {cache_config.model_cfg.num_hidden_layers}" + f" --head_dim {cache_config.model_cfg.head_dim}" + f" --kv_num_head {kv_num_head}" + f" --mp_num {tensor_parallel_size}" @@ -248,7 +248,7 @@ def launch_cache_messager(self, + f" --device_id {int(device_ids[i])}" + f" --rank {i}" + f" --splitwise_role {self.splitwise_role}" - + f" --num_layers {cache_config.model_cfg.num_hidden_layers}" + + f" --num_hidden_layers {cache_config.model_cfg.num_hidden_layers}" + f" --head_dim {cache_config.model_cfg.head_dim}" + f" --kv_num_head {kv_num_head}" + f" --mp_num {tensor_parallel_size}" diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index d5982c687ce..29693803dc1 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -746,10 +746,6 @@ def insert_tasks(self, tasks, current_id=-1, allocated=False): """ Insert tasks to engine. """ - for task in tasks: - start_span_request("DEQUEUE", task, trace.SpanKind.CONSUMER) - if task.sampling_params.bad_words is not None: - task.sampling_params.update_from_tokenizer(self.data_processor.tokenizer) # TODO 返回至 scheduler if allocated: current_tasks = [] @@ -775,6 +771,12 @@ def insert_tasks(self, tasks, current_id=-1, allocated=False): current_tasks.append(cur_task) self.engine_worker_queue.put_tasks((current_tasks, self.resource_manager.real_bsz)) return True + + + for task in tasks: + start_span_request("DEQUEUE", task, trace.SpanKind.CONSUMER) + if task.sampling_params.bad_words is not None: + task.sampling_params.update_from_tokenizer(self.data_processor.tokenizer) self.resource_manager.check_and_free_block_tables() From 57370226203417d1d2b9da82e4927db72671f445 Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Thu, 31 Jul 2025 20:29:22 +0800 Subject: [PATCH 09/13] pre commit format --- fastdeploy/cache_manager/cache_messager.py | 12 +++++------- .../cache_manager/cache_transfer_manager.py | 13 +++++-------- .../cache_manager/prefix_cache_manager.py | 18 ++++-------------- fastdeploy/engine/engine.py | 1 - fastdeploy/spec_decode/mtp.py | 2 +- 5 files changed, 15 insertions(+), 31 deletions(-) diff --git a/fastdeploy/cache_manager/cache_messager.py b/fastdeploy/cache_manager/cache_messager.py index 11853823ca1..0e83f5a6f2d 100644 --- a/fastdeploy/cache_manager/cache_messager.py +++ b/fastdeploy/cache_manager/cache_messager.py @@ -14,22 +14,20 @@ # limitations under the License. """ +import argparse +import json import math -import threading import time -import argparse + import numpy as np import paddle -import json -from fastdeploy.config import SpeculativeConfig from fastdeploy.cache_manager.transfer_factory import IPCCommManager, RDMACommManager +from fastdeploy.config import SpeculativeConfig from fastdeploy.inter_communicator import EngineWorkerQueue, IPCSignal -from fastdeploy.utils import get_logger from fastdeploy.model_executor.ops.gpu import set_data_ipc - def parse_args(): """ 从命令行解析参数 @@ -82,6 +80,7 @@ def parse_args(): args = parser.parse_args() return args + class CacheMessager: """ CacheMessager is used to send the cache data between the engine worker and the cache server. @@ -445,4 +444,3 @@ def main(): print("create cache messager...") print(f"{args}") main() - diff --git a/fastdeploy/cache_manager/cache_transfer_manager.py b/fastdeploy/cache_manager/cache_transfer_manager.py index 2e2d8b1bde8..c9f062201d4 100644 --- a/fastdeploy/cache_manager/cache_transfer_manager.py +++ b/fastdeploy/cache_manager/cache_transfer_manager.py @@ -30,7 +30,6 @@ cuda_host_alloc, share_external_data, swap_cache_all_layers, - ) from fastdeploy.utils import get_logger @@ -124,11 +123,11 @@ def __init__(self, args): cache_type = args.cache_dtype cache_shape = [ - args.num_gpu_blocks, - args.kv_num_head, - args.block_size, - args.head_dim, - ] + args.num_gpu_blocks, + args.kv_num_head, + args.block_size, + args.head_dim, + ] for i in range(args.num_hidden_layers + self.num_extra_layers): num_gpu_blocks = args.num_gpu_blocks if i < args.num_hidden_layers else self.num_extra_layer_gpu_blocks @@ -162,8 +161,6 @@ def __init__(self, args): ) self.v_dst_ptrs.append(self.cpu_cache_kvs[f"value_caches_{i}_rank{rank}"]) - - cache_task_broadcast_data = np.zeros(shape=[1], dtype=np.int32) self.cache_task_broadcast_signal = IPCSignal( name="cache_task_broadcast_signal", diff --git a/fastdeploy/cache_manager/prefix_cache_manager.py b/fastdeploy/cache_manager/prefix_cache_manager.py index 6bbf5511b39..dd101802338 100644 --- a/fastdeploy/cache_manager/prefix_cache_manager.py +++ b/fastdeploy/cache_manager/prefix_cache_manager.py @@ -140,7 +140,6 @@ def launch_cache_manager( filename = "cache_transfer_manager.py" py_path = os.path.join(current_dir_path, filename) - cache_messager_processes = [] if self.splitwise_role != "mixed": cache_messager_processes = self.launch_cache_messager( @@ -165,11 +164,11 @@ def launch_cache_manager( else: kv_num_head = cache_config.model_cfg.num_attention_heads // tensor_parallel_size - log_dir = envs.FD_LOG_DIR cache_manager_processes = [] for i in range(tensor_parallel_size): - launch_cmd = (f" {sys.executable} {py_path}" + launch_cmd = ( + f" {sys.executable} {py_path}" + f" --device_id {int(device_ids[i])}" + f" --rank {i}" + f" --num_hidden_layers {cache_config.model_cfg.num_hidden_layers}" @@ -203,15 +202,8 @@ def launch_cache_manager( cache_manager_processes.extend(cache_messager_processes) return cache_manager_processes - - - def launch_cache_messager(self, - cache_config, - tensor_parallel_size, - device_ids, - pod_ip, - engine_worker_queue_port, - pid_suffix + def launch_cache_messager( + self, cache_config, tensor_parallel_size, device_ids, pod_ip, engine_worker_queue_port, pid_suffix ): """ launch_cache_messager function used to initialize the cache messager. @@ -277,8 +269,6 @@ def launch_cache_messager(self, cache_messager_processes = None return cache_messager_processes - - def update_cache_config(self, cache_config): """ update cache config diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index 29693803dc1..689b4eaf55a 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -771,7 +771,6 @@ def insert_tasks(self, tasks, current_id=-1, allocated=False): current_tasks.append(cur_task) self.engine_worker_queue.put_tasks((current_tasks, self.resource_manager.real_bsz)) return True - for task in tasks: start_span_request("DEQUEUE", task, trace.SpanKind.CONSUMER) diff --git a/fastdeploy/spec_decode/mtp.py b/fastdeploy/spec_decode/mtp.py index c9d0d52d849..e9c1e63a457 100644 --- a/fastdeploy/spec_decode/mtp.py +++ b/fastdeploy/spec_decode/mtp.py @@ -37,8 +37,8 @@ eagle_get_self_hidden_states, mtp_save_first_token, mtp_step_paddle, + set_data_ipc, share_external_data, - set_data_ipc ) from fastdeploy.model_executor.pre_and_post_process import pre_process, rebuild_padding From 857b5cdd3c49f16dd7a4f0d857e4151f323e941e Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Thu, 31 Jul 2025 20:36:13 +0800 Subject: [PATCH 10/13] pre commit format --- fastdeploy/cache_manager/cache_messager.py | 38 +++++++++++----------- fastdeploy/worker/gpu_model_runner.py | 10 +++--- 2 files changed, 25 insertions(+), 23 deletions(-) diff --git a/fastdeploy/cache_manager/cache_messager.py b/fastdeploy/cache_manager/cache_messager.py index 0e83f5a6f2d..4f87e0b2a8c 100644 --- a/fastdeploy/cache_manager/cache_messager.py +++ b/fastdeploy/cache_manager/cache_messager.py @@ -136,7 +136,7 @@ def __init__( ) transfer_protocol = transfer_protocol.split(",") - print(f"splitwise role: {splitwise_role}, {transfer_protocol}" f"rank: {rank}") + logger.info(f"splitwise role: {splitwise_role}, {transfer_protocol}" f"rank: {rank}") # 1. initialize the cache_k_ptr_list and cache_v_ptr_list self.num_hidden_layers = num_hidden_layers @@ -161,7 +161,7 @@ def __init__( block_bytes = math.prod(cache_shape[1:]) if key_cache.dtype == paddle.bfloat16: block_bytes *= 2 - print( + logger.info( f"layers {num_hidden_layers} cache_shape: {cache_shape}, max_block_num: {max_block_num}, " f"block_bytes: {block_bytes}, dtype: {key_cache.dtype}" ) @@ -177,10 +177,10 @@ def __init__( cache_v, ) local_device_id = int(str(cache_k[0].place)[-2]) - print(f"done create ipc_comm with local_device_id:{local_device_id}, ") + logger.info(f"done create ipc_comm with local_device_id:{local_device_id}, ") elif protocol == "rdma": - print(f"splitwise_role rdma: {self.splitwise_role}, rank: {self.rank}, gpu_id: {gpu_id}") + logger.info(f"splitwise_role rdma: {self.splitwise_role}, rank: {self.rank}, gpu_id: {gpu_id}") self.messager[protocol] = RDMACommManager( splitwise_role, @@ -196,7 +196,7 @@ def __init__( self.gpu_id = gpu_id self.cache_info = dict() - print(f"cache messager init finished, use {transfer_protocol}") + logger.info(f"cache messager init finished, use {transfer_protocol}") def prefill_layerwise_send_cache_thread(self): """ @@ -248,7 +248,7 @@ def prefill_layerwise_send_cache_thread(self): cache_info = self.engine_worker_queue.get_cache_info() if cache_info: - print(f"cache info {cache_info}") + logger.info(f"cache info {cache_info}") for info in cache_info: if info["request_id"] in self.cache_info: self.cache_info[info["request_id"]].update(info) @@ -260,7 +260,7 @@ def prefill_layerwise_send_cache_thread(self): current_info["src_block_ids"] = current_src_blocks current_info["current_layer_ids"] = 0 current_info["status"] = "init" - print(f"start cache_infos: {current_info}") + logger.info(f"start cache_infos: {current_info}") self.cache_info[info["request_id"]] = current_info self.last_step_idx = min(self.last_step_idx, current_info["current_id"]) else: @@ -278,7 +278,7 @@ def prefill_layerwise_send_cache_thread(self): if not self.cache_info: time.sleep(0.001) continue - print(f"prefilled_layer_idx: {prefilled_layer_idx}, prefilled_step_idx: {prefilled_step_idx}") + logger.info(f"prefilled_layer_idx: {prefilled_layer_idx}, prefilled_step_idx: {prefilled_step_idx}") for req_id, item in list(self.cache_info.items()): if "status" not in item: continue @@ -295,7 +295,7 @@ def prefill_layerwise_send_cache_thread(self): target_id = int(item["rdma_ports"][self.rank]) status = self.messager[current_transfer_protocol].connect(target_ip, target_id) if not status: - print(f"connect to {target_ip}:{target_id} failed") + logger.info(f"connect to {target_ip}:{target_id} failed") item["status"] = "error" self.engine_worker_queue.finish_request_barrier.wait() if self.rank == 0: @@ -325,7 +325,7 @@ def prefill_layerwise_send_cache_thread(self): self.engine_worker_queue.finish_request_barrier.wait() if self.rank == 0: self.engine_worker_queue.put_finished_req([(item["request_id"], "write cache error")]) - print( + logger.info( f"write cache failed, layer_idx: {layer_idx}, " f"req_id: {item['request_id']}, dest_ip: {target_ip}" ) @@ -336,7 +336,7 @@ def prefill_layerwise_send_cache_thread(self): block_num = len(src_block_ids) avg_time_per_block = cost_time * 1000 / block_num # ms send_cache_speed = block_num * self.block_bytes / 1073741824 / cost_time # GB/s - print( + logger.info( f"finish write cache for a layer, {item['request_id']}, {layer_idx}" f" {current_transfer_protocol}" f"block_num: {block_num}, send_cache_speed(GB/s): {round(send_cache_speed, 5)}," @@ -346,18 +346,18 @@ def prefill_layerwise_send_cache_thread(self): if item["layer_idx"] == self.num_hidden_layers: if item["transfer_protocol"] == "ipc": self.messager["ipc"].write_block_by_sync(target_id) - print(f"finish write cache {item['request_id']}") + logger.info(f"finish write cache {item['request_id']}") self.engine_worker_queue.finish_request_barrier.wait() if self.rank == 0: self.engine_worker_queue.put_finished_req([(item["request_id"], "finished")]) - print(f"put write cache {item['request_id']}") + logger.info(f"put write cache {item['request_id']}") del self.cache_info[req_id] self.last_step_idx = prefilled_step_idx self.last_layer_idx = prefilled_layer_idx except Exception as e: - print(f"prefill layerwise send cache thread has exception: {e}") + logger.info(f"prefill layerwise send cache thread has exception: {e}") def main(): @@ -407,9 +407,9 @@ def main(): f"value_caches_{i}_rank{rank}.device{device}", ) cache_kv_size_byte = sum([tmp.numel() * 1 for key, tmp in gpu_cache_kvs.items()]) - print(f"device :{device}") - print(f"cache_kv_size_byte : {cache_kv_size_byte}") - print(f"done init cache (full) gmem alloc : {paddle.device.cuda.memory_allocated()}") + logger.info(f"device :{device}") + logger.info(f"cache_kv_size_byte : {cache_kv_size_byte}") + logger.info(f"done init cache (full) gmem alloc : {paddle.device.cuda.memory_allocated()}") cache_messager = CacheMessager( splitwise_role=args.splitwise_role, @@ -441,6 +441,6 @@ def main(): args = parse_args() - print("create cache messager...") - print(f"{args}") + logger.info("create cache messager...") + logger.info(f"{args}") main() diff --git a/fastdeploy/worker/gpu_model_runner.py b/fastdeploy/worker/gpu_model_runner.py index c037926634e..59f44edb540 100644 --- a/fastdeploy/worker/gpu_model_runner.py +++ b/fastdeploy/worker/gpu_model_runner.py @@ -43,9 +43,9 @@ from fastdeploy.model_executor.model_loader import get_model_loader from fastdeploy.model_executor.ops.gpu import ( recover_decode_task, + set_data_ipc, set_value_by_flags_and_idx, share_external_data, - set_data_ipc ) from fastdeploy.model_executor.pre_and_post_process import ( post_process, @@ -1148,6 +1148,8 @@ def _update_chunked_prefill(self, tasks): if task.chunk_idx > len(task.prefill_chunk_info): continue self.restore_chunked_prefill_request[task.request_id] = task + if len(self.restore_chunked_prefill_request) > 0: + self.share_inputs["not_need_stop"][0] = True for id, task in list(self.restore_chunked_prefill_request.items()): idx = task.idx @@ -1192,7 +1194,7 @@ def _update_chunked_prefill(self, tasks): self.share_inputs["seq_lens_encoder"][idx : idx + 1] = token_chunk_size self.share_inputs["prompt_lens"][idx : idx + 1] += token_chunk_size self.share_inputs["step_idx"][idx : idx + 1] = 0 - + self.share_inputs["stop_flags"][idx : idx + 1] = False if self.speculative_decoding and self.proposer.is_chunk_prefill_enabled(): self.proposer.update_task_chunk_prefill(task) task.chunk_idx += 1 @@ -1517,12 +1519,12 @@ def cal_theortical_kvcache(self): hidden_dim = self.model_config.head_dim * self.model_config.kv_num_heads # NOTE(liuzichang): Implement multi-layer MTP architecture in the future - num_layers = ( + num_hidden_layers = ( self.model_config.num_hidden_layers + self.speculative_config.num_gpu_block_expand_ratio if self.speculative_method in ["mtp"] else self.model_config.num_hidden_layers ) - required_memory = byte_of_dtype * 2 * (self.cache_config.block_size * hidden_dim) * num_layers # k + v + required_memory = byte_of_dtype * 2 * (self.cache_config.block_size * hidden_dim) * num_hidden_layers # k + v return required_memory def not_need_stop(self) -> bool: From 199521a6e21ee73cabd68c2201be4af0441076c6 Mon Sep 17 00:00:00 2001 From: ltd0924 Date: Fri, 1 Aug 2025 11:07:42 +0800 Subject: [PATCH 11/13] pre commit format --- fastdeploy/cache_manager/cache_messager.py | 2 ++ fastdeploy/engine/engine.py | 11 ++++++----- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/fastdeploy/cache_manager/cache_messager.py b/fastdeploy/cache_manager/cache_messager.py index 4f87e0b2a8c..77f8f3c84a9 100644 --- a/fastdeploy/cache_manager/cache_messager.py +++ b/fastdeploy/cache_manager/cache_messager.py @@ -26,6 +26,7 @@ from fastdeploy.config import SpeculativeConfig from fastdeploy.inter_communicator import EngineWorkerQueue, IPCSignal from fastdeploy.model_executor.ops.gpu import set_data_ipc +from fastdeploy.utils import get_logger def parse_args(): @@ -440,6 +441,7 @@ def main(): if __name__ == "__main__": args = parse_args() + logger = get_logger("cache_messager", "cache_messager.log") logger.info("create cache messager...") logger.info(f"{args}") diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index 689b4eaf55a..e23e2d41240 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -963,14 +963,17 @@ def _exit_sub_services(self): self.running = False if hasattr(self, "cache_manager_processes"): - self.resource_manager.cache_manager.shm_cache_task_flag_broadcast.clear() - self.resource_manager.cache_manager.cache_ready_signal.clear() for p in self.cache_manager_processes: llm_logger.info(f"Killing cache manager process {p.pid}") try: os.killpg(p.pid, signal.SIGTERM) except Exception as e: print(f"Error extracting file: {e}") + if hasattr(self.resource_manager.cache_manager, "cache_ready_signal"): + self.resource_manager.cache_manager.cache_ready_signal.clear() + self.resource_manager.cache_manager.shm_cache_task_flag_broadcast.clear() + if hasattr(self, "zmq_server") and self.zmq_server is not None: + self.zmq_server.close() self.worker_ready_signal.clear() self.exist_task_signal.clear() self.exist_swapped_task_signal.clear() @@ -985,12 +988,10 @@ def _exit_sub_services(self): except Exception as e: print(f"Error extracting sub services: {e}") - self.engine_worker_queue.cleanup() - if hasattr(self, "zmq_server") and self.zmq_server is not None: - self.zmq_server.close() if hasattr(self, "dp_processed"): for p in self.dp_processed: p.join() + self.engine_worker_queue_server.cleanup() def _setting_environ_variables(self): """ From 9f9971844f3249cd48e7882c036553c8c952e026 Mon Sep 17 00:00:00 2001 From: chenjian <1435317881@qq.com> Date: Mon, 4 Aug 2025 20:32:41 +0800 Subject: [PATCH 12/13] [Feature] Support ep pd with external module (#3194) * Support external module * Support external module * Support external module * Support external module * refactor code to make it more clear * refactor code to make it more clear * refactor code to make it more clear * refactor code to make it more clear * fix according to review * fix according to review * fix according to review * fix according to review * fix according to review * fix according to review * fix bug * fix bug * fix bug * merge --------- Co-authored-by: root --- fastdeploy/cache_manager/cache_messager.py | 33 ++- fastdeploy/engine/args_utils.py | 1 + fastdeploy/engine/engine.py | 118 +++++--- fastdeploy/engine/expert_service.py | 28 +- fastdeploy/engine/request.py | 3 + fastdeploy/entrypoints/engine_client.py | 4 +- fastdeploy/envs.py | 8 + fastdeploy/inter_communicator/__init__.py | 5 +- .../inter_communicator/engine_worker_queue.py | 67 +++++ fastdeploy/inter_communicator/zmq_client.py | 188 +++--------- fastdeploy/inter_communicator/zmq_server.py | 273 ++++++++++++++++++ fastdeploy/output/token_processor.py | 12 +- fastdeploy/scheduler/config.py | 68 ++++- fastdeploy/scheduler/dp_scheduler.py | 179 ++++++++++++ .../splitwise/internal_adapter_utils.py | 107 +++++++ 15 files changed, 876 insertions(+), 218 deletions(-) create mode 100644 fastdeploy/inter_communicator/zmq_server.py create mode 100644 fastdeploy/scheduler/dp_scheduler.py create mode 100644 fastdeploy/splitwise/internal_adapter_utils.py diff --git a/fastdeploy/cache_manager/cache_messager.py b/fastdeploy/cache_manager/cache_messager.py index 456ba1c3422..1cc8f4d3120 100644 --- a/fastdeploy/cache_manager/cache_messager.py +++ b/fastdeploy/cache_manager/cache_messager.py @@ -142,12 +142,16 @@ def __init__( self.gpu_id = gpu_id self.cache_info = dict() - self.dp_rank_id = self.rank + local_data_parallel_id * self.nranks + self.rank_id = self.rank + local_data_parallel_id * self.nranks # align with engine worker rank (paddle.distributed.launch) layerwise_send_cache_thread = threading.Thread(target=self._prefill_layerwise_send_cache_thread) layerwise_send_cache_thread.daemon = True layerwise_send_cache_thread.start() + connect_rdma_thread = threading.Thread(target=self._handle_connect_task) + connect_rdma_thread.daemon = True + connect_rdma_thread.start() + logger.info(f"cache messager init finished, use {transfer_protocol}") def _prefill_layerwise_send_cache_thread(self): @@ -160,14 +164,14 @@ def _prefill_layerwise_send_cache_thread(self): prefilled_layer_idx_data = np.zeros(shape=[1], dtype=np.int32) try: step_shm_value = IPCSignal( - name=f"splitwise_complete_prefilled_step_{self.dp_rank_id}", + name=f"splitwise_complete_prefilled_step_{self.rank_id}", array=prefilled_step_idx_data, dtype=np.int32, suffix=self.gpu_id, create=True, ) layer_shm_value = IPCSignal( - name=f"splitwise_complete_prefilled_layer_{self.dp_rank_id}", + name=f"splitwise_complete_prefilled_layer_{self.rank_id}", array=prefilled_layer_idx_data, dtype=np.int32, suffix=self.gpu_id, @@ -175,14 +179,14 @@ def _prefill_layerwise_send_cache_thread(self): ) except: step_shm_value = IPCSignal( - name=f"splitwise_complete_prefilled_step_{self.dp_rank_id}", + name=f"splitwise_complete_prefilled_step_{self.rank_id}", array=prefilled_step_idx_data, dtype=np.int32, suffix=self.gpu_id, create=False, ) layer_shm_value = IPCSignal( - name=f"splitwise_complete_prefilled_layer_{self.dp_rank_id}", + name=f"splitwise_complete_prefilled_layer_{self.rank_id}", array=prefilled_layer_idx_data, dtype=np.int32, suffix=self.gpu_id, @@ -310,3 +314,22 @@ def _prefill_layerwise_send_cache_thread(self): except Exception as e: logger.error(f"prefill layerwise send cache thread has exception: {e}") + + def _handle_connect_task(self): + while True: + try: + task = self.engine_worker_queue.get_connect_rdma_task() + if task is None: + time.sleep(0.001) + continue + logger.info(f"_handle_connect_task recv task: {task}") + task_id = task["task_id"] + ip, rdma_port = task["ip"], task["rdma_port"] + status = self.messager["rdma"].connect(ip, rdma_port) + if not status: + response = {"task_id": task_id, "success": False} + else: + response = {"task_id": task_id, "success": True} + self.engine_worker_queue.put_connect_rdma_task_response(response) + except Exception as e: + logger.error(f"handle_connect_task has exception: {e}") diff --git a/fastdeploy/engine/args_utils.py b/fastdeploy/engine/args_utils.py index d8c57ae4507..2be3c787ace 100644 --- a/fastdeploy/engine/args_utils.py +++ b/fastdeploy/engine/args_utils.py @@ -820,6 +820,7 @@ def create_scheduler_config(self) -> SchedulerConfig: "max_num_partial_prefills", "max_long_partial_prefills", "long_prefill_token_threshold", + "splitwise_role" ] all = asdict(self) diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index e7443bc1db7..6ed5505094c 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -47,12 +47,14 @@ EngineCacheQueue, EngineWorkerQueue, IPCSignal, - ZmqClient, + ZmqIpcServer, + ZmqTcpServer, ) from fastdeploy.metrics.metrics import main_process_metrics from fastdeploy.metrics.trace_util import start_span, start_span_request from fastdeploy.model_executor.guided_decoding import schema_checker from fastdeploy.output.token_processor import TokenProcessor, WarmUpTokenProcessor +from fastdeploy.splitwise.internal_adapter_utils import InternalAdapter from fastdeploy.splitwise.splitwise_connector import SplitwiseConnector from fastdeploy.utils import EngineError, console_logger, envs, llm_logger @@ -179,11 +181,64 @@ def start(self, api_server_pid=None): self.data_processor = self.input_processor.create_processor() if api_server_pid is not None: - self.zmq_server = ZmqClient(name=api_server_pid, mode=zmq.PULL) - self.zmq_server.start_server() - self.zmq_server.create_router() + if envs.FD_ENABLE_INTERNAL_ADAPTER: + self.recv_request_server = ZmqTcpServer(port=envs.FD_ZMQ_RECV_REQUEST_SERVER_PORT, mode=zmq.PULL) + self.send_response_server = ZmqTcpServer(port=envs.FD_ZMQ_SEND_RESPONSE_SERVER_PORT, mode=zmq.ROUTER) + self.external_adapter = InternalAdapter( + cfg=self.cfg, engine=self, dp_rank=self.cfg.node_rank * self.cfg.worker_num_per_node + ) + else: + self.recv_request_server = ZmqIpcServer(name=api_server_pid, mode=zmq.PULL) + self.send_response_server = ZmqIpcServer(name=api_server_pid, mode=zmq.ROUTER) time.sleep(3) + self.cfg.init_cache_info() + + role = self.cfg.splitwise_role + host_ip = self.cfg.host_ip + disaggregate = self.cfg.disaggregate_info + request_queues_for_dp_ipc = ( + None # Different dp has its own process, use multiprocessing.Queue to deliver requests for each dp + ) + result_queue_for_dp_ipc = None + if self.cfg.scheduler_config.name == "splitwise": + self.scheduler.start(role, host_ip, disaggregate) + elif self.cfg.scheduler_config.name == "dp": + request_queues_for_dp_ipc = [] + result_queue_for_dp_ipc = multiprocessing.Queue() + for i in range(self.cfg.parallel_config.data_parallel_size): + request_queues_for_dp_ipc.append(multiprocessing.Queue()) + self.scheduler.start( + self.cfg.node_rank * self.cfg.worker_num_per_node, request_queues_for_dp_ipc, result_queue_for_dp_ipc + ) + + time.sleep(1) + + if self.cfg.parallel_config.enable_expert_parallel and self.cfg.parallel_config.data_parallel_size > 1: + self.dp_processed = [] + for i in range( + 1, + self.cfg.parallel_config.data_parallel_size // self.cfg.nnode, + ): + time.sleep(1) + self.dp_processed.append( + multiprocessing.Process( + target=start_expert_service, + args=( + self.cfg, + i + self.cfg.node_rank * self.cfg.worker_num_per_node, + self.ipc_signal_suffix, + request_queues_for_dp_ipc, + result_queue_for_dp_ipc, + ), + ) + ) + llm_logger.info( + f"Engine is initialized successfully with {self.cfg.tensor_parallel_size}" + + f" data parallel id {i}" + ) + self.dp_processed[-1].start() + if self.do_profile == 0 and ( self.cfg.cache_config.enable_prefix_caching or self.cfg.splitwise_role != "mixed" ): @@ -238,44 +293,11 @@ def start(self, api_server_pid=None): # 单机逻辑 self.engine_worker_queue.available_prefill_instances.put(1) self.split_mode_get_tasks() - if self.cfg.scheduler_config.name == "splitwise": + if self.cfg.scheduler_config.name == "splitwise" or self.cfg.scheduler_config.name == "dp": self.splitwise_receive_thread = threading.Thread(target=self.split_connector.start_receiver, args=()) self.splitwise_receive_thread.daemon = True self.splitwise_receive_thread.start() - self.cfg.init_cache_info() - - role = self.cfg.splitwise_role - host_ip = self.cfg.host_ip - disaggregate = self.cfg.disaggregate_info - if self.cfg.scheduler_config.name == "splitwise": - self.scheduler.start(role, host_ip, disaggregate) - - time.sleep(1) - - if self.cfg.parallel_config.enable_expert_parallel and self.cfg.parallel_config.data_parallel_size > 1: - self.dp_processed = [] - for i in range( - 1, - self.cfg.parallel_config.data_parallel_size // self.cfg.nnode, - ): - time.sleep(1) - self.dp_processed.append( - multiprocessing.Process( - target=start_expert_service, - args=( - self.cfg, - i + self.cfg.node_rank * self.cfg.worker_num_per_node, - self.ipc_signal_suffix, - ), - ) - ) - llm_logger.info( - f"Engine is initialized successfully with {self.cfg.tensor_parallel_size}" - + f" data parallel id {i}" - ) - self.dp_processed[-1].start() - console_logger.info(f"Worker processes are launched with {time.time() - start_time} seconds.") return True @@ -291,7 +313,7 @@ def _zmq_send_generated_tokens(self): time.sleep(0.005) continue for request_id, contents in results.items(): - self.zmq_server.send_multipart(request_id, contents) + self.send_response_server.send_response(request_id, contents) except Exception as e: llm_logger.error(f"Unexcepted error happend: {e}, {traceback.format_exc()!s}") @@ -415,14 +437,18 @@ def _insert_zmq_task_to_scheduler(self): if self.api_server_pid is None: return + if envs.FD_ENABLE_INTERNAL_ADAPTER: + if self.cfg.splitwise_role == "decode": + return + added_requests: Dict[str, int] = dict() while self.running: try: block = True if len(added_requests) == 0 else False if not self.cfg.enable_mm: - err, data = self.zmq_server.receive_json_once(block) + err, data = self.recv_request_server.receive_json_once(block) else: - err, data = self.zmq_server.receive_pyobj_once(block) + err, data = self.recv_request_server.receive_pyobj_once(block) if err is not None: llm_logger.error("Engine stops inserting zmq task into scheduler, err:{err}") break @@ -470,7 +496,7 @@ def _insert_zmq_task_to_scheduler(self): ) # Since the request is not in scheduler # Send result by zmq directly - self.zmq_server.send_multipart(request_id, error_result) + self.send_response_server.send_response(request_id, error_result) except Exception as e: llm_logger.error( f"Error happend while receving new request from zmq, details={e}, " @@ -989,8 +1015,12 @@ def _exit_sub_services(self): print(f"Error extracting sub services: {e}") self.engine_worker_queue.cleanup() - if hasattr(self, "zmq_server") and self.zmq_server is not None: - self.zmq_server.close() + if hasattr(self, "send_response_server") and self.send_response_server is not None: + self.send_response_server.close() + if hasattr(self, "recv_request_server") and self.recv_request_server is not None: + self.recv_request_server.close() + if hasattr(self, "recv_control_cmd_server") and self.recv_control_cmd_server is not None: + self.recv_control_cmd_server.close() if hasattr(self, "dp_processed"): for p in self.dp_processed: p.join() diff --git a/fastdeploy/engine/expert_service.py b/fastdeploy/engine/expert_service.py index 63b1b15beba..6b3c0147605 100644 --- a/fastdeploy/engine/expert_service.py +++ b/fastdeploy/engine/expert_service.py @@ -29,8 +29,9 @@ from fastdeploy.inter_communicator import EngineWorkerQueue from fastdeploy.metrics.metrics import main_process_metrics from fastdeploy.output.token_processor import TokenProcessor +from fastdeploy.splitwise.internal_adapter_utils import InternalAdapter from fastdeploy.splitwise.splitwise_connector import SplitwiseConnector -from fastdeploy.utils import EngineError, console_logger, llm_logger +from fastdeploy.utils import EngineError, console_logger, envs, llm_logger class ExpertService: @@ -60,7 +61,8 @@ def __init__(self, cfg, local_data_parallel_id): self.scheduler = cfg.scheduler_config.scheduler() - self.scheduler.reset_nodeid(f"{self.scheduler.infer.nodeid}_{local_data_parallel_id!s}") + if self.cfg.scheduler_config.name == "splitwise": + self.scheduler.reset_nodeid(f"{self.scheduler.infer.nodeid}_{local_data_parallel_id!s}") self.cfg.parallel_config.local_data_parallel_id = local_data_parallel_id @@ -111,8 +113,12 @@ def __init__(self, cfg, local_data_parallel_id): ) self._finalizer = weakref.finalize(self, self._exit_sub_services) + if envs.FD_ENABLE_INTERNAL_ADAPTER: + self.external_adapter = InternalAdapter(cfg=self.cfg, engine=self, dp_rank=local_data_parallel_id) - def start(self, ipc_signal_suffix, local_data_parallel_id): + def start( + self, ipc_signal_suffix, local_data_parallel_id, request_queues_for_dp_ipc=None, result_queue_for_dp_ipc=None + ): """ Initializes the engine and starts its sub-services. If `api_server_pid` is defined, will launch a thread @@ -127,7 +133,7 @@ def start(self, ipc_signal_suffix, local_data_parallel_id): cache_config=self.cfg.cache_config, tensor_parallel_size=self.cfg.tensor_parallel_size, device_ids=self.cfg.local_device_ids, - pod_ip=self.cfg.pod_ips[0], + pod_ip=self.cfg.master_ip, engine_worker_queue_port=self.cfg.engine_worker_queue_port, pid_suffix=f"{local_data_parallel_id}_{ipc_signal_suffix}", ) @@ -147,7 +153,11 @@ def start(self, ipc_signal_suffix, local_data_parallel_id): role = self.cfg.splitwise_role host_ip = self.cfg.host_ip disaggregate = self.cfg.disaggregate_info - self.scheduler.start(role, host_ip, disaggregate) + if self.cfg.scheduler_config.name == "dp": + assert (request_queues_for_dp_ipc is not None) and (result_queue_for_dp_ipc is not None) + self.scheduler.start(local_data_parallel_id, request_queues_for_dp_ipc, result_queue_for_dp_ipc) + elif self.cfg.scheduler_config.name == "splitwise": + self.scheduler.start(role, host_ip, disaggregate) self.cfg.print() console_logger.info(f"Worker processes are launched with {time.time() - start_time} seconds.") @@ -356,13 +366,17 @@ def _exit_sub_services(self): self.zmq_server.close() -def start_expert_service(cfg, local_data_parallel_id, ipc_signal_suffix): +def start_expert_service( + cfg, local_data_parallel_id, ipc_signal_suffix, request_queues_for_dp_ipc=None, result_queue_for_dp_ipc=None +): """ Start expert service """ expert_service = ExpertService(cfg, local_data_parallel_id) try: - expert_service.start(ipc_signal_suffix, local_data_parallel_id) + expert_service.start( + ipc_signal_suffix, local_data_parallel_id, request_queues_for_dp_ipc, result_queue_for_dp_ipc + ) expert_service.split_connector.start_receiver() except Exception as e: llm_logger.exception(f"Expert service failed to start: {e}") diff --git a/fastdeploy/engine/request.py b/fastdeploy/engine/request.py index acf717547a7..f88d24152be 100644 --- a/fastdeploy/engine/request.py +++ b/fastdeploy/engine/request.py @@ -71,6 +71,7 @@ def __init__( guided_json_object: Optional[bool] = None, enable_thinking: Optional[bool] = True, trace_carrier: dict = dict(), + dp_rank: Optional[int] = None ) -> None: self.request_id = request_id self.prompt = prompt @@ -119,6 +120,7 @@ def __init__( self.task_type = RequestType.PREFILL self.idx = None self.need_prefill_tokens = self.prompt_token_ids_len + self.dp_rank = dp_rank @classmethod def from_dict(cls, d: dict): @@ -151,6 +153,7 @@ def from_dict(cls, d: dict): guided_json_object=d.get("guided_json_object", None), enable_thinking=d.get("enable_thinking", True), trace_carrier=d.get("trace_carrier", {}), + dp_rank=d.get("dp_rank", None) ) @property diff --git a/fastdeploy/entrypoints/engine_client.py b/fastdeploy/entrypoints/engine_client.py index 09d6e8ff9fe..fad81d6af97 100644 --- a/fastdeploy/entrypoints/engine_client.py +++ b/fastdeploy/entrypoints/engine_client.py @@ -21,7 +21,7 @@ from fastdeploy.engine.config import ModelConfig from fastdeploy.input.preprocess import InputPreprocessor -from fastdeploy.inter_communicator import IPCSignal, ZmqClient +from fastdeploy.inter_communicator import IPCSignal, ZmqIpcClient from fastdeploy.metrics.work_metrics import work_process_metrics from fastdeploy.multimodal.registry import MultimodalRegistry from fastdeploy.platforms import current_platform @@ -90,7 +90,7 @@ def create_zmq_client(self, model, mode): """ Create a ZMQ client. """ - self.zmq_client = ZmqClient(model, mode) + self.zmq_client = ZmqIpcClient(model, mode) self.zmq_client.connect() def format_and_add_data(self, prompts: dict): diff --git a/fastdeploy/envs.py b/fastdeploy/envs.py index 3da7af75cbb..dd95de1df5d 100644 --- a/fastdeploy/envs.py +++ b/fastdeploy/envs.py @@ -80,6 +80,14 @@ "EXPORTER_OTLP_HEADERS": lambda: os.getenv("EXPORTER_OTLP_HEADERS"), # enable kv cache block scheduler v1 (no need for kv_cache_ratio) "ENABLE_V1_KVCACHE_SCHEDULER": lambda: int(os.getenv("ENABLE_V1_KVCACHE_SCHEDULER", "0")), + # enable internal module to access LLMEngine. + "FD_ENABLE_INTERNAL_ADAPTER": lambda: int(os.getenv("FD_ENABLE_INTERNAL_ADAPTER", "0")), + # LLMEngine recieve requests port, used when FD_ENABLE_INTERNAL_ADAPTER=1 + "FD_ZMQ_RECV_REQUEST_SERVER_PORT": lambda: os.getenv("FD_ZMQ_RECV_REQUEST_SERVER_PORT", "8200"), + # LLMEngine send response port, used when FD_ENABLE_INTERNAL_ADAPTER=1 + "FD_ZMQ_SEND_RESPONSE_SERVER_PORT": lambda: os.getenv("FD_ZMQ_SEND_RESPONSE_SERVER_PORT", "8201"), + # LLMEngine recieve control command port, used when FD_ENABLE_INTERNAL_ADAPTER=1 + "FD_ZMQ_CONTROL_CMD_SERVER_PORTS": lambda: os.getenv("FD_ZMQ_CONTROL_CMD_SERVER_PORTS", "8202"), # Whether to use PLUGINS. "FD_PLUGINS": lambda: None if "FD_PLUGINS" not in os.environ else os.environ["FD_PLUGINS"].split(","), } diff --git a/fastdeploy/inter_communicator/__init__.py b/fastdeploy/inter_communicator/__init__.py index 0c1cc0d9fc9..ea08af31a40 100644 --- a/fastdeploy/inter_communicator/__init__.py +++ b/fastdeploy/inter_communicator/__init__.py @@ -17,6 +17,7 @@ from .engine_cache_queue import EngineCacheQueue from .engine_worker_queue import EngineWorkerQueue from .ipc_signal import IPCSignal -from .zmq_client import ZmqClient +from .zmq_client import ZmqIpcClient +from .zmq_server import ZmqIpcServer, ZmqTcpServer -__all__ = ["ZmqClient", "IPCSignal", "EngineWorkerQueue", "EngineCacheQueue"] +__all__ = ["ZmqIpcClient", "IPCSignal", "EngineWorkerQueue", "EngineCacheQueue", "ZmqTcpServer", "ZmqIpcServer"] diff --git a/fastdeploy/inter_communicator/engine_worker_queue.py b/fastdeploy/inter_communicator/engine_worker_queue.py index da88265a266..e216f430d2c 100644 --- a/fastdeploy/inter_communicator/engine_worker_queue.py +++ b/fastdeploy/inter_communicator/engine_worker_queue.py @@ -85,12 +85,15 @@ class QueueManager(BaseManager): ] self.finished_req_queue = [Queue() for _ in range(self.local_data_parallel_size)] self.cache_infos_init: List[List[Any]] = [list() for _ in range(self.local_data_parallel_size)] + self.connect_rdma_tasks_list = [list() for _ in range(self.local_data_parallel_size)] + self.connect_rdma_tasks_response_list = [list() for _ in range(self.local_data_parallel_size)] self.client_read_info_flag_init: List[List[int]] = [ [1] * self.num_client for _ in range(self.local_data_parallel_size) ] self.lock_info_init: List[threading.Lock] = [ threading.Lock() for _ in range(self.local_data_parallel_size) ] + self.connect_task_lock_init: List[threading.Lock] = [threading.Lock() for _ in range(self.local_data_parallel_size)] self.finish_request_barrier = [ threading.Barrier(self.num_client) for _ in range(self.local_data_parallel_size) @@ -112,11 +115,26 @@ class QueueManager(BaseManager): callable=lambda idx: self.lock_init[idx], proxytype=AcquirerProxy, ) + QueueManager.register( + "get_connect_task_lock", + callable=lambda idx: self.connect_task_lock_init[idx], + proxytype=AcquirerProxy, + ) QueueManager.register( "get_read_finish_flag", callable=lambda idx: self.read_finish_flag_init[idx], proxytype=ValueProxy, ) + QueueManager.register( + "get_connect_rdma_tasks", + callable=lambda idx: self.connect_rdma_tasks_list[idx], + proxytype=ListProxy + ) + QueueManager.register( + "get_connect_rdma_tasks_responses", + callable=lambda idx: self.connect_rdma_tasks_response_list[idx], + proxytype=ListProxy + ) QueueManager.register( "get_connected_client_counter", callable=lambda idx: self.connected_client_counter_init[idx], @@ -180,6 +198,9 @@ class QueueManager(BaseManager): QueueManager.register("get_disaggregate_requests") QueueManager.register("get_available_prefill_instances") QueueManager.register("get_finish_request_barrier") + QueueManager.register("get_connect_rdma_tasks") + QueueManager.register("get_connect_rdma_tasks_responses") + QueueManager.register("get_connect_task_lock") self.manager = QueueManager(address=self.address, authkey=self.authkey) self._connect_with_retry() @@ -200,6 +221,13 @@ class QueueManager(BaseManager): self.available_prefill_instances = self.manager.get_available_prefill_instances() self.finish_request_barrier = self.manager.get_finish_request_barrier(self.local_data_parallel_id) self.finished_req_queue = self.manager.get_finish_request_queue(self.local_data_parallel_id) + # p/d互联 + self.connect_rdma_task_queue = self.manager.get_connect_rdma_tasks(self.local_data_parallel_id) + self.connect_rdma_task_response_queue = self.manager.get_connect_rdma_tasks_responses( + self.local_data_parallel_id + ) + self.connect_task_lock = self.manager.get_connect_task_lock(self.local_data_parallel_id) + assert self.num_client == len(self.client_read_flag) if is_server: @@ -280,6 +308,45 @@ def num_tasks(self) -> int: total_num: int = len(self.tasks) self.lock.release() return total_num + + def put_connect_rdma_task(self, connect_rdma_task): + self.connect_task_lock.acquire() + self.connect_rdma_task_queue.append(connect_rdma_task) + self.connect_task_lock.release() + + def get_connect_rdma_task(self): + result = None + self.connect_task_lock.acquire() + if len(self.connect_rdma_task_queue) == 0: + self.connect_task_lock.release() + return result + try: + result = self.connect_rdma_task_queue.pop(0) + except Exception as e: + llm_logger.info(f"get_connect_rdma_task got exception: {e}") + finally: + self.connect_task_lock.release() + return result + + def put_connect_rdma_task_response(self, connect_rdma_task_response): + self.connect_task_lock.acquire() + self.connect_rdma_task_response_queue.append(connect_rdma_task_response) + self.connect_task_lock.release() + + def get_connect_rdma_task_response(self): + result = None + self.connect_task_lock.acquire() + if len(self.connect_rdma_task_response_queue) == 0: + self.connect_task_lock.release() + return result + try: + result = self.connect_rdma_task_response_queue.pop(0) + except Exception as e: + llm_logger.info(f"get_connect_rdma_task_response got exception: {e}") + finally: + self.connect_task_lock.release() + return result + def get_prefill_instances(self): """ diff --git a/fastdeploy/inter_communicator/zmq_client.py b/fastdeploy/inter_communicator/zmq_client.py index 05e55929dda..13242f2a204 100644 --- a/fastdeploy/inter_communicator/zmq_client.py +++ b/fastdeploy/inter_communicator/zmq_client.py @@ -14,200 +14,78 @@ # limitations under the License. """ -import os -import threading -import time +from abc import ABC, abstractmethod -import msgpack import zmq -from fastdeploy import envs -from fastdeploy.utils import llm_logger - -class ZmqClient: +class ZmqClientBase(ABC): """ - ZmqClient is a class that provides a client-side interface for sending and receiving messages using ZeroMQ. + ZmqClientBase is a base class that provides a client-side interface for sending and receiving messages using ZeroMQ. """ - def __init__(self, name, mode): - self.context = zmq.Context() - self.socket = self.context.socket(mode) - self.file_name = f"/dev/shm/{name}.socket" - self.router_path = f"/dev/shm/router_{name}.ipc" + def __init__(self): + pass - self.ZMQ_SNDHWM = int(envs.FD_ZMQ_SNDHWM) - self.aggregate_send = envs.FD_USE_AGGREGATE_SEND + @abstractmethod + def _create_socket(self): + """Abstract method to create and return a ZeroMQ socket.""" + pass - self.mutex = threading.Lock() - self.req_dict = dict() - self.router = None - self.poller = None - self.running = True + def _ensure_socket(self): + """Ensure the socket is created before use.""" + if self.socket is None: + self.socket = self._create_socket() + @abstractmethod def connect(self): """ Connect to the server using the file name specified in the constructor. """ - self.socket.connect(f"ipc://{self.file_name}") - - def start_server(self): - """ - Start the server using the file name specified in the constructor. - """ - self.socket.setsockopt(zmq.SNDHWM, self.ZMQ_SNDHWM) - self.socket.setsockopt(zmq.SNDTIMEO, -1) - self.socket.bind(f"ipc://{self.file_name}") - self.poller = zmq.Poller() - self.poller.register(self.socket, zmq.POLLIN) - - def create_router(self): - """ - Create a ROUTER socket and bind it to the specified router path. - """ - self.router = self.context.socket(zmq.ROUTER) - self.router.setsockopt(zmq.SNDHWM, self.ZMQ_SNDHWM) - self.router.setsockopt(zmq.SNDTIMEO, -1) - self.router.bind(f"ipc://{self.router_path}") + pass def send_json(self, data): """ Send a JSON-serializable object over the socket. """ + self._ensure_socket() self.socket.send_json(data) def recv_json(self): """ Receive a JSON-serializable object from the socket. """ + self._ensure_socket() return self.socket.recv_json() def send_pyobj(self, data): """ Send a Pickle-serializable object over the socket. """ + self._ensure_socket() self.socket.send_pyobj(data) def recv_pyobj(self): """ Receive a Pickle-serializable object from the socket. """ + self._ensure_socket() return self.socket.recv_pyobj() - def pack_aggregated_data(self, data): - """ - Aggregate multiple responses into one and send them to the client. - """ - result = data[0] - if len(data) > 1: - for response in data[1:]: - result.add(response) - result = msgpack.packb([result.to_dict()]) - return result - - def send_multipart(self, req_id, data): - """ - Send a multipart message to the router socket. - """ - if self.router is None: - raise RuntimeError("Router socket not created. Call create_router() first.") - - while self.running: - with self.mutex: - if req_id not in self.req_dict: - try: - client, _, request_id = self.router.recv_multipart(flags=zmq.NOBLOCK) - req_id_str = request_id.decode("utf-8") - self.req_dict[req_id_str] = client - except zmq.Again: - time.sleep(0.001) - continue - else: - break - - try: - start_send = time.time() - if self.aggregate_send: - result = self.pack_aggregated_data(data) - else: - result = msgpack.packb([response.to_dict() for response in data]) - self.router.send_multipart([self.req_dict[req_id], b"", result]) - llm_logger.debug(f"send_multipart result: {req_id} len {len(data)} elapse: {time.time()-start_send}") - - except Exception as e: - llm_logger.error(f"Send result to zmq client failed: {e}") - - if data[-1].finished: - with self.mutex: - self.req_dict.pop(req_id, None) - llm_logger.info(f"send_multipart finished, req_id: {req_id}") - - def receive_json_once(self, block=False): - """ - Receive a single message from the socket. - """ - if self.socket is None or self.socket.closed: - return "zmp socket has closed", None - try: - flags = zmq.NOBLOCK if not block else 0 - return None, self.socket.recv_json(flags=flags) - except zmq.Again: - return None, None - except Exception as e: - self.close() - llm_logger.warning(f"{e}") - return str(e), None - - def receive_pyobj_once(self, block=False): - """ - Receive a single message from the socket. - """ - if self.socket is None or self.socket.closed: - return "zmp socket has closed", None - try: - flags = zmq.NOBLOCK if not block else 0 - return None, self.socket.recv_pyobj(flags=flags) - except zmq.Again: - return None, None - except Exception as e: - self.close() - llm_logger.warning(f"{e}") - return str(e), None - - def _clear_ipc(self, name): - """ - Remove the IPC file with the given name. - """ - if os.path.exists(name): - try: - os.remove(name) - except OSError as e: - llm_logger.warning(f"Failed to remove IPC file {name} - {e}") - - def close(self): - """ - Close the socket and context, and remove the IPC files. - """ - if not self.running: - return - - self.running = False - llm_logger.info("Closing ZMQ connection...") - try: - if hasattr(self, "socket") and not self.socket.closed: - self.socket.close() - - if self.router is not None and not self.router.closed: - self.router.close() - if not self.context.closed: - self.context.term() +class ZmqIpcClient(ZmqClientBase): + def __init__(self, name, mode): + self.name = name + self.mode = mode + self.file_name = f"/dev/shm/{name}.socket" + self.context = zmq.Context() + self.socket = self.context.socket(self.mode) - self._clear_ipc(self.file_name) - self._clear_ipc(self.router_path) - except Exception as e: - llm_logger.warning(f"Failed to close ZMQ connection - {e}") - return + def _create_socket(self): + """create and return a ZeroMQ socket.""" + self.context = zmq.Context() + return self.context.socket(self.mode) - def __exit__(self, exc_type, exc_val, exc_tb): - self.close() + def connect(self): + self._ensure_socket() + self.socket.connect(f"ipc://{self.file_name}") diff --git a/fastdeploy/inter_communicator/zmq_server.py b/fastdeploy/inter_communicator/zmq_server.py new file mode 100644 index 00000000000..f4ee8be313d --- /dev/null +++ b/fastdeploy/inter_communicator/zmq_server.py @@ -0,0 +1,273 @@ +""" +# Copyright (c) 2025 PaddlePaddle Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License" +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" + +import os +import threading +import time +from abc import ABC, abstractmethod + +import msgpack +import zmq + +from fastdeploy import envs +from fastdeploy.utils import llm_logger + + +class ZmqServerBase(ABC): + """ + ZmqServerBase + """ + + def __init__(self): + pass + + @abstractmethod + def _create_socket(self): + """Abstract method to create and return a ZeroMQ socket.""" + pass + + def _ensure_socket(self): + """Ensure the socket is created before use.""" + if self.socket is None: + self.socket = self._create_socket() + + def pack_aggregated_data(self, data): + """ + Aggregate multiple responses into one and send them to the client. + """ + result = data[0] + if len(data) > 1: + for response in data[1:]: + result.add(response) + result = msgpack.packb([result.to_dict()]) + return result + + def receive_json_once(self, block=False): + """ + Receive a single message from the socket. + """ + self._ensure_socket() + if self.socket is None or self.socket.closed: + return "zmp socket has closed", None + try: + flags = zmq.NOBLOCK if not block else 0 + return None, self.socket.recv_json(flags=flags) + except zmq.Again: + return None, None + except Exception as e: + self.close() + llm_logger.warning(f"{e}") + return str(e), None + + def receive_pyobj_once(self, block=False): + """ + Receive a single message from the socket. + """ + self._ensure_socket() + if self.socket is None or self.socket.closed: + return "zmp socket has closed", None + try: + flags = zmq.NOBLOCK if not block else 0 + return None, self.socket.recv_pyobj(flags=flags) + except zmq.Again: + return None, None + except Exception as e: + self.close() + llm_logger.warning(f"{e}") + return str(e), None + + def send_response(self, req_id, data): + """ + Send generated token result to client. + """ + self._ensure_socket() + if self.socket is None: + raise RuntimeError("Router socket not created. Call create_router() first.") + + while self.running: + with self.mutex: + if req_id not in self.req_dict: + try: + client, _, request_id = self.socket.recv_multipart(flags=zmq.NOBLOCK) + req_id_str = request_id.decode("utf-8") + self.req_dict[req_id_str] = client + except zmq.Again: + time.sleep(0.001) + continue + else: + break + + try: + start_send = time.time() + if self.aggregate_send: + result = self.pack_aggregated_data(data) + else: + result = msgpack.packb([response.to_dict() for response in data]) + self.socket.send_multipart([self.req_dict[req_id], b"", result]) + llm_logger.debug(f"send_multipart result: {req_id} len {len(data)} elapse: {time.time()-start_send}") + + except Exception as e: + llm_logger.error(f"Send result to zmq client failed: {e}") + + if data[-1].finished: + with self.mutex: + self.req_dict.pop(req_id, None) + llm_logger.info(f"send_multipart finished, req_id: {req_id}") + + @abstractmethod + def close(self): + pass + + def __exit__(self, exc_type, exc_val, exc_tb): + self.close() + + +class ZmqIpcServer(ZmqServerBase): + """ + ZmqIpcServer, used when FD_ENABLE_INTERNAL_ADAPTER=0 + """ + + def __init__(self, name, mode): + self.name = name + self.mode = mode + if mode == zmq.PULL: + self.file_name = f"/dev/shm/{name}.socket" + elif mode == zmq.ROUTER: + self.file_name = f"/dev/shm/router_{name}.ipc" + self.ZMQ_SNDHWM = int(envs.FD_ZMQ_SNDHWM) + self.aggregate_send = envs.FD_USE_AGGREGATE_SEND + self.mutex = threading.Lock() + self.req_dict = dict() + self.running = True + self.context = zmq.Context() + self._create_socket() + + def _create_socket(self): + """create and return a ZeroMQ socket.""" + self.socket = self.context.socket(self.mode) + self.socket.setsockopt(zmq.SNDHWM, self.ZMQ_SNDHWM) + self.socket.setsockopt(zmq.SNDTIMEO, -1) + self.socket.bind(f"ipc://{self.file_name}") + return self.socket + + def _clear_ipc(self, name): + """ + Remove the IPC file with the given name. + """ + if os.path.exists(name): + try: + os.remove(name) + except OSError as e: + llm_logger.warning(f"Failed to remove IPC file {name} - {e}") + + def close(self): + """ + Close the socket and context, and remove the IPC files. + """ + if not self.running: + return + + self.running = False + llm_logger.info("Closing ZMQ connection...") + try: + if self.socket is not None and not self.socket.closed: + self.socket.close() + if not self.context.closed: + self.context.term() + self._clear_ipc(self.file_name) + except Exception as e: + llm_logger.warning(f"Failed to close ZMQ connection - {e}") + return + + +class ZmqTcpServer(ZmqServerBase): + """ + ZmqTcpServer, used when FD_ENABLE_INTERNAL_ADAPTER=1 + """ + + def __init__(self, port, mode): + self.mode = mode + self.port = port + self.ZMQ_SNDHWM = int(envs.FD_ZMQ_SNDHWM) + self.aggregate_send = envs.FD_USE_AGGREGATE_SEND + + self.mutex = threading.Lock() + self.req_dict = dict() + self.running = True + self.context = zmq.Context() + self._create_socket() + + def _create_socket(self): + """create and return a ZeroMQ socket.""" + self.socket = self.context.socket(self.mode) + self.socket.setsockopt(zmq.SNDHWM, self.ZMQ_SNDHWM) + self.socket.setsockopt(zmq.SNDTIMEO, -1) + self.socket.bind(f"tcp://*:{self.port}") + return self.socket + + def recv_control_cmd(self): + """ + Recieve control command from client + """ + self._ensure_socket() + while self.running: + try: + client, _, task_data = self.socket.recv_multipart(flags=zmq.NOBLOCK) + task = msgpack.unpackb(task_data) + task_id_str = task["task_id"] + except zmq.Again: + time.sleep(0.001) + continue + with self.mutex: + self.req_dict[task_id_str] = client + return task + + def response_for_control_cmd(self, task_id, result): + """ + Send command result back to client. + """ + self._ensure_socket() + if self.socket is None: + raise RuntimeError("Router socket not created.") + try: + result = msgpack.packb(result) + self.socket.send_multipart([self.req_dict[task_id], b"", result]) + + except Exception as e: + llm_logger.error(f"Send result to zmq client failed: {e}") + + with self.mutex: + self.req_dict.pop(task_id, None) + llm_logger.info(f"response control cmd finished, task_id: {task_id}") + + def close(self): + """ + Close the socket and context. + """ + if not self.running: + return + + self.running = False + llm_logger.info("Closing ZMQ connection...") + try: + if self.socket is not None and not self.socket.closed: + self.socket.close() + if not self.context.closed: + self.context.term() + + except Exception as e: + llm_logger.warning(f"Failed to close ZMQ connection - {e}") + return diff --git a/fastdeploy/output/token_processor.py b/fastdeploy/output/token_processor.py index 27fda998700..3938db36246 100644 --- a/fastdeploy/output/token_processor.py +++ b/fastdeploy/output/token_processor.py @@ -412,7 +412,11 @@ def _process_sampling_with_logprob_batch_output(self): self._record_completion_metrics(task, current_time) self._recycle_resources(task_id, i, task, result, is_prefill) break - if not is_prefill or self.cfg.scheduler_config.name == "splitwise": + if ( + not is_prefill + or self.cfg.scheduler_config.name == "splitwise" + or self.cfg.scheduler_config.name == "dp" + ): batch_result.append(result) self.postprocess(batch_result) @@ -531,7 +535,11 @@ def _process_batch_output(self): self._record_completion_metrics(task, current_time) self._recycle_resources(task_id, i, task, result, is_prefill) break - if not is_prefill or self.cfg.scheduler_config.name == "splitwise": + if ( + not is_prefill + or self.cfg.scheduler_config.name == "splitwise" + or self.cfg.scheduler_config.name == "dp" + ): batch_result.append(result) self.postprocess(batch_result) diff --git a/fastdeploy/scheduler/config.py b/fastdeploy/scheduler/config.py index cd0a72af1a2..c831b0f44ad 100644 --- a/fastdeploy/scheduler/config.py +++ b/fastdeploy/scheduler/config.py @@ -18,6 +18,7 @@ from fastdeploy.utils import llm_logger +from .dp_scheduler import DPScheduler from .global_scheduler import GlobalScheduler from .local_scheduler import LocalScheduler from .splitwise_scheduler import SplitWiseScheduler, SplitWiseSchedulerConfig @@ -89,6 +90,57 @@ def print(self): llm_logger.info("=============================================================") +class DPLocalSchedulerConfig(LocalSchedulerConfig): + """ + Configuration class for DPLocalScheduler. + + Attributes: + max_size: Maximum number of concurrent requests (-1 for unlimited) + ttl: Time-to-live in seconds for request expiration + """ + + def __init__( + self, + max_size: int = -1, + ttl: int = 900, + max_model_len: int = 8192, + enable_chunked_prefill: bool = False, + max_num_partial_prefills: int = 1, + max_long_partial_prefills: int = 1, + long_prefill_token_threshold: int = 0, + splitwise_role: str = "prefill", + **kwargs, + ): + """ + Initialize LocalScheduler configuration. + + Args: + max_size: Maximum concurrent requests (-1 for unlimited, 0 for disabled) + ttl: Time-to-live in seconds for request expiration (default 900s) + max_model_len: Maximum model context length in tokens + enable_chunked_prefill: Whether to enable chunked prefill processing + max_num_partial_prefills: Max partial prefill operations allowed + max_long_partial_prefills: Max long-running partial prefill ops + long_prefill_token_threshold: Token count threshold for long prefill + **kwargs: Additional unused arguments (for forward compatibility) + + Note: + - If long_prefill_token_threshold is 0, it's auto-calculated as 4% of max_model_len + - See LocalScheduler class for implementation details + """ + self.max_size = max_size + self.ttl = ttl + + self.max_model_len = max_model_len + self.enable_chunked_prefill = enable_chunked_prefill + self.max_num_partial_prefills = max_num_partial_prefills + self.max_long_partial_prefills = max_long_partial_prefills + self.long_prefill_token_threshold = long_prefill_token_threshold + if self.long_prefill_token_threshold == 0: + self.long_prefill_token_threshold = int(self.max_model_len * 0.04) + self.splitwise_role = splitwise_role + + class GlobalSchedulerConfig: """ Configuration class for GlobalScheduler (Redis-based). @@ -229,6 +281,9 @@ def __init__(self, name="local", **kwargs): if name == "splitwise": self.config = SplitWiseSchedulerConfig(**kwargs) + if name == "dp": + self.config = DPLocalSchedulerConfig(**kwargs) + def check(self): """ Validate the configuration. @@ -236,7 +291,7 @@ def check(self): Raises: Exception: If invalid scheduler type is specified """ - if self.name not in ["local", "global", "splitwise"]: + if self.name not in ["local", "global", "splitwise", "dp"]: raise Exception(f"Unknown scheduler type {self.name}") self.config.check() @@ -274,6 +329,17 @@ def scheduler(self): if self.name == "splitwise": return SplitWiseScheduler(self.config) + if self.name == "dp": + return DPScheduler( + max_size=self.config.max_size, + ttl=self.config.ttl, + enable_chunked_prefill=self.config.enable_chunked_prefill, + max_num_partial_prefills=self.config.max_num_partial_prefills, + max_long_partial_prefills=self.config.max_long_partial_prefills, + long_prefill_token_threshold=self.config.long_prefill_token_threshold, + splitwise_role=self.config.splitwise_role, + ) + return LocalScheduler( max_size=self.config.max_size, ttl=self.config.ttl, diff --git a/fastdeploy/scheduler/dp_scheduler.py b/fastdeploy/scheduler/dp_scheduler.py new file mode 100644 index 00000000000..d55a687905e --- /dev/null +++ b/fastdeploy/scheduler/dp_scheduler.py @@ -0,0 +1,179 @@ +""" +# Copyright (c) 2025 PaddlePaddle Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" + +import threading +import time +from multiprocessing import Queue +from typing import Dict, List, Optional + +from fastdeploy.engine.request import Request, RequestOutput +from fastdeploy.scheduler.data import ScheduledResponse +from fastdeploy.scheduler.local_scheduler import LocalScheduler +from fastdeploy.utils import scheduler_logger + + +class DPLocalScheduler(LocalScheduler): + def __init__( + self, + max_size: int, + ttl: int, + enable_chunked_prefill: bool, + max_num_partial_prefills: int, + max_long_partial_prefills: int, + long_prefill_token_threshold: int, + splitwise_role: str = "prefill", + ): + super().__init__( + max_size, + ttl, + enable_chunked_prefill, + max_num_partial_prefills, + max_long_partial_prefills, + long_prefill_token_threshold, + ) + self.splitwise_role = splitwise_role + + def put_results(self, results: List[RequestOutput]): + """ + Add processing results back to the scheduler. + Args: + results: List of RequestOutput objects containing results + """ + responses: List[ScheduledResponse] = [ScheduledResponse(result) for result in results] + + finished_responses = [response.request_id for response in responses if response.finished] + if len(finished_responses) > 0: + scheduler_logger.info(f"Scheduler has received some finished responses: {finished_responses}") + + with self.mutex: + for response in responses: + if response.request_id not in self.responses: + self.responses[response.request_id] = [response] + continue + self.responses[response.request_id].append(response) + self.responses_not_empty.notify_all() + + def _recycle(self, request_id: Optional[str] = None): + """ + Clean up expired or completed requests to free memory. + Args: + request_id: Optional specific request ID to remove. + If None, removes all expired requests. + """ + if request_id is not None: + self.requests.pop(request_id, None) + self.responses.pop(request_id, None) + if self.splitwise_role == "decode": + return + self.ids.pop(self.ids.index(request_id)) + self.ids_read_cursor -= 1 + return + + if self.max_size <= 0: + return + + if len(self.requests) <= self.max_size: + return + + now = time.time() + expired_ids = [] + for request_id in self.ids: + request = self.requests[request_id] + if now - request.schedule_time < self.ttl: + break + expired_ids.append(request.request_id) + + for i, expired_id in enumerate(expired_ids): + self.requests.pop(expired_id, None) + self.responses.pop(expired_id, None) + self.ids.pop(i) + + if len(expired_ids) > 0: + if len(expired_ids) - 1 >= self.ids_read_cursor: + self.ids_read_cursor = 0 + else: + self.ids_read_cursor -= len(expired_ids) + + +class DPScheduler: + def __init__( + self, + max_size: int, + ttl: int, + enable_chunked_prefill: bool, + max_num_partial_prefills: int, + max_long_partial_prefills: int, + long_prefill_token_threshold: int, + splitwise_role: str = "prefill", + ): + self._scheduler = DPLocalScheduler( + max_size, + ttl, + enable_chunked_prefill, + max_num_partial_prefills, + max_long_partial_prefills, + long_prefill_token_threshold, + splitwise_role, + ) + + def start(self, dp_rank: int, request_queues: List[Queue], result_queue: Queue): + self.dp_rank = dp_rank + self.request_queues = request_queues + self.result_queue = result_queue + threading.Thread(target=self._put_requests_to_local).start() + threading.Thread(target=self._get_response_from_local).start() + + def put_requests(self, requests: List[Dict]): + results = [] + for request in requests: + if not hasattr(request, "dp_rank"): + raise ValueError(f"Request object is missing the 'dp_rank' attribute: {request}") + self.request_queues[request.dp_rank].put(request) + results.append((request.request_id, None)) + return results + + def _put_requests_to_local(self): + while True: + request = self.request_queues[self.dp_rank].get() + self._scheduler.put_requests([request]) + + def _get_response_from_local(self): + while True: + results = self._scheduler.get_results() + if len(results) == 0: + continue + self.result_queue.put(results) + + def get_requests( + self, + available_blocks, + block_size, + reserved_output_blocks, + max_num_batched_tokens, + batch=1, + ) -> List[Request]: + return self._scheduler.get_requests( + available_blocks, block_size, reserved_output_blocks, max_num_batched_tokens, batch + ) + + def get_unhandled_request_num(self): + return len(self._scheduler.requests) + + def put_results(self, results: List[RequestOutput]): + self._scheduler.put_results(results) + + def get_results(self) -> Dict[str, List[RequestOutput]]: + return self.result_queue.get() diff --git a/fastdeploy/splitwise/internal_adapter_utils.py b/fastdeploy/splitwise/internal_adapter_utils.py new file mode 100644 index 00000000000..db3ea520d4a --- /dev/null +++ b/fastdeploy/splitwise/internal_adapter_utils.py @@ -0,0 +1,107 @@ +""" +# Copyright (c) 2025 PaddlePaddle Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License" +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" + +import threading +import time +import traceback + +# **Note**: Just for internal use +import zmq + +from fastdeploy.inter_communicator import ZmqTcpServer +from fastdeploy.metrics.metrics import get_filtered_metrics, main_process_metrics +from fastdeploy.utils import envs, get_logger + +logger = get_logger("internal_adapter_utils", "internal_adapter_utils.log") + + +class InternalAdapter: + def __init__(self, cfg, engine, dp_rank): + self.cfg = cfg + self.engine = engine + self.dp_rank = dp_rank + recv_control_cmd_ports = envs.FD_ZMQ_CONTROL_CMD_SERVER_PORTS.split(",") + self.recv_control_cmd_server = ZmqTcpServer(port=recv_control_cmd_ports[dp_rank], mode=zmq.ROUTER) + self.recv_external_instruct_thread = threading.Thread( + target=self._recv_external_module_control_instruct, daemon=True + ) + self.recv_external_instruct_thread.start() + self.response_external_instruct_thread = threading.Thread( + target=self._response_external_module_control_instruct, daemon=True + ) + self.response_external_instruct_thread.start() + + def _get_current_server_info(self): + """ + Get resources information + """ + available_batch_size = min(self.cfg.max_prefill_batch, self.engine.resource_manager.available_batch()) + + available_block_num = self.engine.resource_manager.available_block_num() + server_info = { + "splitwise_role": self.cfg.splitwise_role, + "block_size": int(self.cfg.cache_config.block_size), + "block_num": int(available_block_num), + "dec_token_num": int(self.cfg.cache_config.dec_token_num), + "available_resource": 1.0 * available_block_num / self.cfg.cache_config.total_block_num, + "max_batch_size": int(available_batch_size), + "max_input_token_num": self.cfg.max_num_batched_tokens, + "unhandled_request_num": self.engine.scheduler.get_unhandled_request_num(), + } + return server_info + + def _recv_external_module_control_instruct(self): + """ + Receive a multipart message from the control cmd socket. + """ + while True: + try: + task = self.recv_control_cmd_server.recv_control_cmd() + logger.info(f"Recieve control task: {task}") + task_id_str = task["task_id"] + if task["cmd"] == "get_payload": + payload_info = self._get_current_server_info() + result = {"task_id": task_id_str, "result": payload_info} + logger.info(f"Response for task: {task_id_str}") + self.recv_control_cmd_server.response_for_control_cmd(task_id_str, result) + + elif task["cmd"] == "get_metrics": + metrics_text = get_filtered_metrics( + [], + extra_register_func=lambda reg: main_process_metrics.register_all(reg, workers=1), + ) + result = {"task_id": task_id_str, "result": metrics_text} + logger.info(f"Response for task: {task_id_str}") + self.recv_control_cmd_server.response_for_control_cmd(task_id_str, result) + elif task["cmd"] == "connect_rdma": + self.engine.engine_worker_queue.put_connect_rdma_task(task) + + except Exception as e: + logger.error(f"handle_control_cmd got error: {e}, {traceback.format_exc()!s}") + + def _response_external_module_control_instruct(self): + while True: + try: + result_data = self.engine.engine_worker_queue.get_connect_rdma_task_response() + if result_data: + task_id_str = result_data["task_id"] + result = {"task_id": task_id_str, "result": result_data} + logger.info(f"Response for task: {task_id_str}") + self.recv_control_cmd_server.response_for_control_cmd(task_id_str, result) + else: + time.sleep(0.001) + except Exception as e: + logger.error(f"_handle_connect_rdma_results got error: {e}, {traceback.format_exc() !s}") From 5474baea0934ba47e0c6436bcd15bb9b87a1fe97 Mon Sep 17 00:00:00 2001 From: ltd0924 <32387785+ltd0924@users.noreply.github.com> Date: Tue, 5 Aug 2025 14:13:10 +0800 Subject: [PATCH 13/13] Update cache_messager.py --- fastdeploy/cache_manager/cache_messager.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/fastdeploy/cache_manager/cache_messager.py b/fastdeploy/cache_manager/cache_messager.py index 80191c58c0b..547fcfd607a 100644 --- a/fastdeploy/cache_manager/cache_messager.py +++ b/fastdeploy/cache_manager/cache_messager.py @@ -18,7 +18,7 @@ import json import math import time - +import threading import numpy as np import paddle @@ -469,4 +469,4 @@ def main(): logger.info("create cache messager...") logger.info(f"{args}") - main() \ No newline at end of file + main()