From d927adf8a6f7d593d4bb2de3ad8f054adac72db0 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Sun, 4 Jan 2026 19:38:30 +0800 Subject: [PATCH 01/25] support dynamic run_control_request through zmq from apiserver to common_engine --- fastdeploy/engine/common_engine.py | 71 +++++++++++++++++++++- fastdeploy/engine/request.py | 94 ++++++++++++++++++++++++++++++ 2 files changed, 164 insertions(+), 1 deletion(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index e756a8e004d..56c5a9ae398 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -38,7 +38,7 @@ from tqdm import tqdm import fastdeploy.metrics.trace as tracing -from fastdeploy.engine.request import Request, RequestOutput, RequestType +from fastdeploy.engine.request import Request, ControlRequest, RequestOutput, RequestType from fastdeploy.engine.resource_manager import ResourceManager from fastdeploy.engine.sched.resource_manager_v1 import ResourceManagerV1 from fastdeploy.eplb.utils import init_eplb_signals @@ -1064,6 +1064,17 @@ def _insert_zmq_task_to_scheduler(self): self.llm_logger.error(f"Engine stops inserting zmq task into scheduler, err:{err}") break + if ControlRequest.is_control_request(data): + try: + control_req = ControlRequest.from_dict(data) + self.run_control_method(control_req) + except Exception as e: + self.llm_logger.error( + f"Failed to process control request {data.get('request_id')}: " + f"{e}, {traceback.format_exc()}" + ) + continue + request, insert_task = None, [] results: List[Tuple[str, Optional[str]]] = list() if data: @@ -1115,6 +1126,64 @@ def _insert_zmq_task_to_scheduler(self): f"traceback={traceback.format_exc()}" ) + def run_control_method(self, control_req: ControlRequest): + """ + Execute control methods for engine management using dynamic method invocation. + + Args: + control_req: ControlRequest instance containing method name and arguments + + Usage: + - Control request with method "get_metrics" will call self._control_get_metrics(args) + - Method names are automatically mapped to handler methods with prefix "_control_" + - If no handler exists, returns error with available methods + """ + method = control_req.get_method() + args = control_req.get_args() + request_id = control_req.request_id + + try: + self.llm_logger.info(f"Processing control request {request_id}: {method}") + + # Dynamically map method name to handler method + handler_name = f"_control_{method}" + handler = getattr(self, handler_name, None) + if handler is None or not callable(handler): + error_msg = f"Unknown control method: {method}" + self.llm_logger.error(errmsg) + self._send_error_response(request_id, 400, error_msg) + return + + # Dynamically call the handler method with provided arguments + error_code, error_msg = handler(args) + if error_code == 0: + self.llm_logger.error(f"Control method {method} failed: {error_msg}") + self._send_error_response(request_id, error_msg, error_code) + return + + self.llm_logger.info(f"Control method {method} success.") + succ_result = RequestOutput(request_id=request_id, finished=True) + self.send_response_server.send_response(request_id, [succ_result]) + + except Exception as e: + error_msg = f"Control method {method} failed: {str(e)}" + self.llm_logger.error(f"{error_msg}\n{traceback.format_exc()}") + self._send_error_response(request_id, 500, error_msg) + + def _control_pause(self, args: dict) -> dict: + """暂停请求生成 + + Args: + args: 控制参数字典,暂停相关的配置参数 + + Returns: + tuple: (error_code, error_msg) 元组 + - error_code: 错误代码,0表示成功,非0表示失败 + - error_msg: 错误信息,成功时为空字符串 + """ + self.llm_logger.info(f"Pause Request Generation") + return 0, "" + def _send_error_response(self, request_id, error_msg, error_code: int = 500): self.llm_logger.error( f"Send error response to client, request_id: {request_id}, error_msg: {error_msg}, error_code: {error_code}" diff --git a/fastdeploy/engine/request.py b/fastdeploy/engine/request.py index 6a7b665de53..1e150121850 100644 --- a/fastdeploy/engine/request.py +++ b/fastdeploy/engine/request.py @@ -376,6 +376,100 @@ def __repr__(self) -> str: except Exception as e: return f"" +class ControlRequest: + """A generic control request that supports method and args for control operations. + + This request type is used for system-level control operations rather than + typical inference requests. It enables dynamic control of engine behavior, + resource management, and system configuration via a flexible method-args interface. + """ + + def __init__( + self, + request_id: str, + method: str, + args: Optional[Dict[str, Any]] = None, + ) -> None: + """ + Args: + request_id: Unique identifier for the control request. + method: The control method to execute (e.g., "reset_scheduler", "get_metrics"). + args: Optional arguments for the control method. + """ + self.request_id = request_id + self.method = method + self.args = args or {} + + @classmethod + def from_dict(cls, d: dict): + """Create ControlRequest instance from dictionary.""" + return cls( + request_id=d["request_id"], + method=d["method"], + args=d.get("args", {}) + ) + + def to_dict(self) -> dict: + """Convert ControlRequest into a serializable dict.""" + return { + "request_id": self.request_id, + "method": self.method, + "args": self.args + } + + def __repr__(self) -> str: + """Provide a clean representation of the control request.""" + try: + if not envs.FD_DEBUG: + return f"ControlRequest(request_id={self.request_id}, method={self.method})" + else: + return ( + f"ControlRequest(" + f"request_id={self.request_id}, " + f"method={self.method}, " + f"args={self.args}" + f")" + ) + except Exception as e: + return f"" + + def get_method(self) -> str: + """Get the control method name.""" + return self.method + + def get_args(self) -> Dict[str, Any]: + """Get the control method arguments.""" + return self.args.copy() + + @staticmethod + def is_control_request(d: dict) -> bool: + """ + Check if a dictionary represents a valid ControlRequest. + + Args: + d: Dictionary to check + + Returns: + bool: True if the dictionary contains the required fields for a ControlRequest + """ + + # Check if all required fields are present and have correct types + if not isinstance(d, dict): + return False + + # Check field types + if "request_id" not in d or not isinstance(d.get("request_id"), str): + return False + + if "method" not in d or not isinstance(d.get("method"), str): + return False + + # Args is optional, but if present should be a dict + if "args" in d and not isinstance(d["args"], dict): + return False + + return True + @dataclass(slots=True) class CompletionOutput: From b1e1f0e22e3e82be47021885a8e5109df4434be7 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Mon, 5 Jan 2026 15:15:10 +0800 Subject: [PATCH 02/25] support pause/resume/is_paused/update_weights in apiserver->common_engine by common run_control_method --- fastdeploy/engine/common_engine.py | 62 ++++++++++++---- fastdeploy/engine/request.py | 81 +++++++++++++++++++-- fastdeploy/entrypoints/engine_client.py | 20 +++++ fastdeploy/entrypoints/openai/api_server.py | 29 ++++++++ 4 files changed, 168 insertions(+), 24 deletions(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index 56c5a9ae398..7fbca77dfcc 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -38,7 +38,7 @@ from tqdm import tqdm import fastdeploy.metrics.trace as tracing -from fastdeploy.engine.request import Request, ControlRequest, RequestOutput, RequestType +from fastdeploy.engine.request import Request, ControlRequest, ControlResponse, RequestOutput, RequestType from fastdeploy.engine.resource_manager import ResourceManager from fastdeploy.engine.sched.resource_manager_v1 import ResourceManagerV1 from fastdeploy.eplb.utils import init_eplb_signals @@ -1145,32 +1145,26 @@ def run_control_method(self, control_req: ControlRequest): try: self.llm_logger.info(f"Processing control request {request_id}: {method}") - # Dynamically map method name to handler method handler_name = f"_control_{method}" handler = getattr(self, handler_name, None) if handler is None or not callable(handler): - error_msg = f"Unknown control method: {method}" - self.llm_logger.error(errmsg) - self._send_error_response(request_id, 400, error_msg) + error_result = ControlResponse(request_id, 400, f"unknown control method:{method}") + self.llm_logger.error(str(error_result)) + self.send_response_server.send_response(request_id, [error_result]) return - # Dynamically call the handler method with provided arguments - error_code, error_msg = handler(args) - if error_code == 0: - self.llm_logger.error(f"Control method {method} failed: {error_msg}") - self._send_error_response(request_id, error_msg, error_code) - return - + result = handler(args) self.llm_logger.info(f"Control method {method} success.") - succ_result = RequestOutput(request_id=request_id, finished=True) + succ_result = ControlResponse(request_id, 200, "Success", result) self.send_response_server.send_response(request_id, [succ_result]) except Exception as e: error_msg = f"Control method {method} failed: {str(e)}" self.llm_logger.error(f"{error_msg}\n{traceback.format_exc()}") - self._send_error_response(request_id, 500, error_msg) + error_result = ControlResponse(request_id, 500, error_msg) + self.send_response_server.send_response(request_id, [error_result]) - def _control_pause(self, args: dict) -> dict: + def _control_pause(self, args: dict) -> dict | None: """暂停请求生成 Args: @@ -1182,7 +1176,43 @@ def _control_pause(self, args: dict) -> dict: - error_msg: 错误信息,成功时为空字符串 """ self.llm_logger.info(f"Pause Request Generation") - return 0, "" + return None + + def _control_resume(self, args: dict) -> dict | None: + """恢复暂停的请求生成 + + Args: + args: 控制参数字典,恢复生成相关的配置参数 + + Returns: + dict | None: 返回结果字典或None,包含恢复操作的状态信息 + """ + self.llm_logger.info(f"Resume Request Generation") + return None + + def _control_is_paused(self, args: dict) -> bool: + """检查是否暂停了请求生成 + + Args: + args: 控制参数字典,检查是否暂停相关的配置参数 + + Returns: + bool: 是否暂停了请求生成 + """ + self.llm_logger.info(f"Check if Request Generation is Paused") + return {"is_paused": self.is_paused} + + def _control_update_weights(self, args: dict) -> dict | None: + """更新模型权重 + + Args: + args: 控制参数字典,更新权重相关的配置参数 + + Returns: + dict | None: 返回结果字典或None,包含更新权重的操作结果信息 + """ + self.llm_logger.info(f"Update Model Weights") + return None def _send_error_response(self, request_id, error_msg, error_code: int = 500): self.llm_logger.error( diff --git a/fastdeploy/engine/request.py b/fastdeploy/engine/request.py index 1e150121850..3d16dc17bfd 100644 --- a/fastdeploy/engine/request.py +++ b/fastdeploy/engine/request.py @@ -21,6 +21,7 @@ from dataclasses import asdict, dataclass, fields from enum import Enum from typing import Any, Dict, Generic, Optional, Union +from fastapi.responses import JSONResponse import numpy as np from typing_extensions import TypeVar @@ -86,7 +87,7 @@ def __init__( guided_json_object: Optional[bool] = None, enable_thinking: Optional[bool] = True, reasoning_max_tokens: Optional[int] = None, - trace_carrier: dict = dict(), + trace_carrier: Optional[Dict[str, Any]] = None, dp_rank: Optional[int] = None, chat_template: Optional[str] = None, image_start: int = 0, @@ -204,13 +205,16 @@ def from_dict(cls, d: dict): # if mm_positions is not of type ImagePosition, convert to ImagePosition try: for i, mm_pos in enumerate(d["multimodal_inputs"]["mm_positions"]): - d["multimodal_inputs"]["mm_positions"][i] = ( - ImagePosition(**mm_pos) if not isinstance(mm_pos, ImagePosition) else mm_pos - ) - except Exception as e: + if not isinstance(mm_pos, ImagePosition): + if not isinstance(mm_pos, dict): + raise ValueError(f"Invalid mm_positions format at index {i}") + d["multimodal_inputs"]["mm_positions"][i] = ImagePosition(**mm_pos) + except (ValueError, TypeError, KeyError) as e: data_processor_logger.error( - f"Convert mm_positions to ImagePosition error: {e}, {str(traceback.format_exc())}" + f"Convert mm_positions to ImagePosition failed - {type(e).__name__}: {e}\n" + f"Input data: {d['multimodal_inputs']['mm_positions']}" ) + raise return cls( request_id=d["request_id"], prompt=d.get("prompt"), @@ -470,6 +474,67 @@ def is_control_request(d: dict) -> bool: return True +class ControlResponse: + """ + Response for control opeartions + """ + def __init__( + self, + request_id: str, + error_code: int = 200, + error_message: Optional[str] = None, + result: Optional[dict] = None, + finished: bool = True) -> None: + self.request_id = request_id + self.finished = finished + self.error_message = error_message + self.result = result + self.error_code = error_code + + def to_dict(self) -> dict: + """Convert ControlResponse into a serializable dict.""" + return { + "request_id": self.request_id, + "finished": self.finished, + "error_code": self.error_code, + "error_message": self.error_message, + "result": self.result + } + + @classmethod + def from_dict(cls, d: dict): + """Create ControlResponse instance from dictionary.""" + return cls( + request_id=d["request_id"], + finished=d.get("finished", True), + error_code=d.get("error_code", 200), + error_message=d.get("error_message"), + result=d.get("result") + ) + + def to_api_json_response(self) -> JSONResponse: + """Convert ControlResponse into a JSONResponse.""" + status = "success" if self.error_code == 200 else "error" + content = { + "request_id": self.request_id, + "status": status, + "error_message": self.error_message, + "result": self.result + } + return JSONResponse(status_code=self.error_code, content=content) + + def __repr__(self) -> str: + """Provide a clean representation of the control response.""" + return ( + f"ControlResponse(" + f"request_id={self.request_id}, " + f"finished={self.finished}, " + f"error_code={self.error_code}, " + f"error_message={self.error_message}, " + f"result={self.result}" + f")" + ) + @dataclass(slots=True) class CompletionOutput: @@ -905,8 +970,8 @@ class PoolingRequestOutput(Generic[_O]): prompt_token_ids: list[int] finished: bool metrics: Optional[RequestMetrics] = (None,) - error_code: Optional[int] = (200,) - error_msg: Optional[str] = (None,) + error_code: Optional[int] = 200 + error_msg: Optional[str] = None def __repr__(self): return ( diff --git a/fastdeploy/entrypoints/engine_client.py b/fastdeploy/entrypoints/engine_client.py index 9babe8fec74..41b10697387 100644 --- a/fastdeploy/entrypoints/engine_client.py +++ b/fastdeploy/entrypoints/engine_client.py @@ -21,6 +21,7 @@ import uuid from copy import copy from http import HTTPStatus +import asyncio import numpy as np from filelock import FileLock @@ -51,6 +52,7 @@ api_server_logger, to_tensor, ) +from fastdeploy.engine.request import ControlRequest, ControlResponse class EngineClient: @@ -512,6 +514,24 @@ def check_health(self, time_interval_threashold=30): return True, "" + async def run_control_method(self, request: ControlRequest): + api_server_logger.info(f"Start Run Control Method: {request}") + self.zmq_client.send_json(request.to_dict()) + request_id = request.request_id + dealer, response_queue = await self.connection_manager.get_connection(request_id) + dealer.write([b"", request_id.encode("utf-8")]) + try: + response = await asyncio.wait_for(response_queue.get(), timeout=600) + print(response) + response = ControlResponse.from_dict(response[0]) + api_server_logger.info(f"End Run Control Method: {response}") + return response + except asyncio.TimeoutError: + error_response = ControlResponse(request_id, 500, "Timeout waiting for control method response") + api_server_logger.error(f"Error Run Control Method: {error_response}") + return error_response + + def is_workers_alive(self): """ Check the health of the model server by checking whether all workers are alive. diff --git a/fastdeploy/entrypoints/openai/api_server.py b/fastdeploy/entrypoints/openai/api_server.py index 13879f72a84..b84efb391d8 100644 --- a/fastdeploy/entrypoints/openai/api_server.py +++ b/fastdeploy/entrypoints/openai/api_server.py @@ -22,6 +22,7 @@ import traceback from collections.abc import AsyncGenerator from contextlib import asynccontextmanager +import uuid import uvicorn import zmq @@ -37,6 +38,7 @@ from fastdeploy.engine.args_utils import EngineArgs from fastdeploy.engine.async_llm import AsyncLLM from fastdeploy.engine.engine import LLMEngine +from fastdeploy.engine.request import ControlRequest from fastdeploy.engine.expert_service import ExpertService from fastdeploy.entrypoints.chat_utils import load_chat_template from fastdeploy.entrypoints.engine_client import EngineClient @@ -365,6 +367,33 @@ def ping(raw_request: Request) -> Response: """Ping check. Endpoint required for SageMaker""" return health(raw_request) +@app.post("/v1/pause") +async def pause(request: Request) -> Response: + request_id = f"control-{uuid.uuid4()}" + control_request = ControlRequest(request_id, "pause") + control_response = await app.state.engine_client.run_control_method(control_request) + return control_response.to_api_json_response() + +@app.post("/v1/resume") +async def resume(request: Request) -> Response: + request_id = f"control-{uuid.uuid4()}" + control_request = ControlRequest(request_id, "resume") + control_response = await app.state.engine_client.run_control_method(control_request) + return control_response.to_api_json_response() + +@app.get("/v1/is_paused") +async def is_paused(request: Request) -> Response: + request_id = f"control-{uuid.uuid4()}" + control_request = ControlRequest(request_id, "is_paused") + control_response = await app.state.engine_client.run_control_method(control_request) + return control_response.to_api_json_response() + +@app.post("/v1/update_weights") +async def update_weights(request: Request) -> Response: + request_id = f"control-{uuid.uuid4()}" + control_request = ControlRequest(request_id, "update_weights") + control_response = await app.state.engine_client.run_control_method(control_request) + return control_response.to_api_json_response() def wrap_streaming_generator(original_generator: AsyncGenerator): """ From c877720ac27204ca3bdd45a9d88660e42ce5d15a Mon Sep 17 00:00:00 2001 From: mitu626 Date: Mon, 5 Jan 2026 15:23:53 +0800 Subject: [PATCH 03/25] change /is_puased from HTTP POST method to GET method --- fastdeploy/entrypoints/openai/api_server.py | 1 + 1 file changed, 1 insertion(+) diff --git a/fastdeploy/entrypoints/openai/api_server.py b/fastdeploy/entrypoints/openai/api_server.py index b84efb391d8..29d6c04b407 100644 --- a/fastdeploy/entrypoints/openai/api_server.py +++ b/fastdeploy/entrypoints/openai/api_server.py @@ -369,6 +369,7 @@ def ping(raw_request: Request) -> Response: @app.post("/v1/pause") async def pause(request: Request) -> Response: + # todo: support wait_for_inflight_requests(default False), clear_cache(default True) arguments request_id = f"control-{uuid.uuid4()}" control_request = ControlRequest(request_id, "pause") control_response = await app.state.engine_client.run_control_method(control_request) From c0bf6653bd22ba4bbb1be0dce0e672f0514fdff9 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 8 Jan 2026 20:45:03 +0800 Subject: [PATCH 04/25] =?UTF-8?q?add=20pause=E3=80=81resume=E3=80=81is=5Fp?= =?UTF-8?q?aused=20implementation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- fastdeploy/engine/common_engine.py | 71 +++++++++++++++++-- .../engine/sched/resource_manager_v1.py | 40 +++++++++++ fastdeploy/output/token_processor.py | 2 +- 3 files changed, 105 insertions(+), 8 deletions(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index 7fbca77dfcc..16a8a219202 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -81,6 +81,9 @@ def __init__(self, cfg, start_queue=True, use_async_llm=False): """ self.cfg = cfg self.use_async_llm = use_async_llm + + self.is_paused = False # pause request generation + self._pause_cond = threading.Condition() if self.cfg.parallel_config.data_parallel_size > 1: self.llm_logger = get_logger( @@ -760,6 +763,8 @@ def _schedule_request_to_worker_v1(self): def _fetch_request(): try: + with self._pause_cond: + self._pause_cond.wait_for(lambda: not self.is_paused) nonlocal is_fetching num_prefill_batch = min( int(self.resource_manager.available_batch()), @@ -923,6 +928,8 @@ def _fetch_request(): is_fetching = False while self.running: + with self._pause_cond: + self._pause_cond.wait_for(lambda: not self.is_paused) try: if self.engine_worker_queue.exist_tasks(): time.sleep(0.001) @@ -1087,6 +1094,11 @@ def _insert_zmq_task_to_scheduler(self): trace_print(LoggingEventName.REQUEST_SCHEDULE_START, data["request_id"], data.get("user", "")) trace_print(LoggingEventName.REQUEST_QUEUE_START, data["request_id"], data.get("user", "")) self.llm_logger.debug(f"Receive request from api server: {request}") + + if self.is_paused: + self.llm_logger.warning(f"Engine is paused, drop request: {request}") + self._send_error_response(request.request_id, "Request is aborted since LLM Engine is paused.") + continue except Exception as e: self.llm_logger.error(f"Receive request error: {e}, {traceback.format_exc()!s}") err_msg = str(e) @@ -1143,7 +1155,7 @@ def run_control_method(self, control_req: ControlRequest): request_id = control_req.request_id try: - self.llm_logger.info(f"Processing control request {request_id}: {method}") + self.llm_logger.info(f"START run control method {request_id}: {method}") handler_name = f"_control_{method}" handler = getattr(self, handler_name, None) @@ -1154,12 +1166,12 @@ def run_control_method(self, control_req: ControlRequest): return result = handler(args) - self.llm_logger.info(f"Control method {method} success.") + self.llm_logger.info(f"SUCCESS run control method {method}.") succ_result = ControlResponse(request_id, 200, "Success", result) self.send_response_server.send_response(request_id, [succ_result]) except Exception as e: - error_msg = f"Control method {method} failed: {str(e)}" + error_msg = f"Failed run control method {method}: {str(e)}" self.llm_logger.error(f"{error_msg}\n{traceback.format_exc()}") error_result = ControlResponse(request_id, 500, error_msg) self.send_response_server.send_response(request_id, [error_result]) @@ -1175,7 +1187,45 @@ def _control_pause(self, args: dict) -> dict | None: - error_code: 错误代码,0表示成功,非0表示失败 - error_msg: 错误信息,成功时为空字符串 """ - self.llm_logger.info(f"Pause Request Generation") + with self._pause_cond: + if self.is_paused: + self.llm_logger.info("Pause Request Generation: already paused.") + self.is_paused = True + + self.llm_logger.info(f"Start Abort Running Requests") + + self.resource_manager.log_status() + # preempted all running reqs. preempted reqs will be append to ResourceManager.waiting queue + timeout, count = 60, 0 + while self.engine_worker_queue.exist_tasks(): + time.sleep(0.001) + count += 1 + if count >= timeout * 1000: + break + if count >= timeout * 1000: + error_msg = f"wait engine_worker_queue tasks empty timeout after {timeout} seconds, worker may Hanged" + self.llm_logger.info(error_msg) + raise Exception(error_msg) + running_reqs = self.resource_manager.preempted_all() + if len(running_reqs) > 0: + self.llm_logger.info(f"Total {len(running_reqs)} requests need to be aborted.") + self.resource_manager.get_real_bsz() + self.engine_worker_queue.put_tasks((running_reqs, self.resource_manager.real_bsz)) + self.resource_manager.wait_worker_inflight_requests_finish(timeout=60) + self.resource_manager.log_status() + self.engine_worker_queue.clear_data() + self.token_processor.clear_data() + #self.resource_manager.clear_data() + self.resource_manager.log_status() + + # abort inflight requests to user + inflight_requests = [req.raw for req in self.scheduler.requests.values()] + self.llm_logger.info(f"Start Abort Inflight Requests, total {len(inflight_requests)} waiting requests") + for req in inflight_requests: + self._send_error_response(req.request_id, "Request is aborted since LLM Engine is paused.") + self.scheduler.reset() + + self.resource_manager.cache_manager.reset() return None def _control_resume(self, args: dict) -> dict | None: @@ -1187,7 +1237,13 @@ def _control_resume(self, args: dict) -> dict | None: Returns: dict | None: 返回结果字典或None,包含恢复操作的状态信息 """ - self.llm_logger.info(f"Resume Request Generation") + self.llm_logger.info(f"START Resume Request Generation") + with self._pause_cond: + if not self.is_paused: + self.llm_logger.info("Resume Request Generation: not paused.") + self.is_paused = False + self._pause_cond.notify_all() + self.llm_logger.info(f"END Resume Request Generation") return None def _control_is_paused(self, args: dict) -> bool: @@ -1199,8 +1255,9 @@ def _control_is_paused(self, args: dict) -> bool: Returns: bool: 是否暂停了请求生成 """ - self.llm_logger.info(f"Check if Request Generation is Paused") - return {"is_paused": self.is_paused} + self.llm_logger.info(f"LLM Engine request generation is paused: {self.is_paused}") + with self._pause_cond: + return {"is_paused": self.is_paused} def _control_update_weights(self, args: dict) -> dict | None: """更新模型权重 diff --git a/fastdeploy/engine/sched/resource_manager_v1.py b/fastdeploy/engine/sched/resource_manager_v1.py index 91c5160a31d..f102a9a5890 100644 --- a/fastdeploy/engine/sched/resource_manager_v1.py +++ b/fastdeploy/engine/sched/resource_manager_v1.py @@ -227,6 +227,7 @@ def _prepare_preempt_task(self, request): def reschedule_preempt_task(self, request_id, process_func=None): with self.lock: + llm_logger.debug(f"reschedule {request_id} into waiting queue") if request_id in self.to_be_rescheduled_request_id_set and request_id in self.requests: request = self.requests[request_id] if process_func is not None: @@ -252,6 +253,34 @@ def _can_preempt(self): return True return False + def preempted_all(self): + with self.lock: + preempted_reqs = [] + for i in range(len(self.running)): + req = self.running.pop() + req.status = RequestStatus.PREEMPTED + req.num_computed_tokens = 0 + #self.tasks_list[req.idx] = None + #self.stop_flags[req.idx] = True + self._free_blocks(req) + req.cached_block_num = 0 + self.to_be_rescheduled_request_id_set.add(req.request_id) + preempted_reqs.append(self._prepare_preempt_task(req)) + return preempted_reqs + + def wait_worker_inflight_requests_finish(self, timeout=60): + count = 0 + while count < timeout * 1000: + running_reqs_count = len(self.to_be_rescheduled_request_id_set) + if running_reqs_count == 0: + break + + count += 1 + time.sleep(0.001) + if count >= timeout * 1000: + llm_logger.info(f"wait_inflight_requests_finish timeout after {timeout} seconds, " + f"still {len(self.to_be_rescheduled_request_id_set)} requests running") + def _trigger_preempt(self, request, num_new_blocks, preempted_reqs, scheduled_reqs): """ If the request cannot be scheduled, preempt the running request one by one until it can be scheduled. Last in, first out. @@ -1235,3 +1264,14 @@ def update_metrics(self): main_process_metrics.gpu_cache_usage_perc.set(self.get_gpu_cache_usage_perc()) main_process_metrics.num_requests_running.set(len(self.running)) main_process_metrics.num_requests_waiting.set(num_tasks - len(self.running)) + + def log_status(self): + llm_logger.info(f"ResourceManagerV1( " + f"waiting={len(self.waiting)}, " + f"running={len(self.running)}, " + f"preempted={len(self.to_be_rescheduled_request_id_set)}, " + f"tasks_list={self.tasks_list}, " + f"stop_flags={self.stop_flags}, " + f"req_dict={self.req_dict}, " + f"requests={self.requests}, " + f")") diff --git a/fastdeploy/output/token_processor.py b/fastdeploy/output/token_processor.py index 578fe583c1c..021bad80cc4 100644 --- a/fastdeploy/output/token_processor.py +++ b/fastdeploy/output/token_processor.py @@ -987,7 +987,7 @@ def clear_data(self): finished=True, metrics=RequestMetrics( arrival_time=time.time(), - request_start_time=task.arrival_time, + request_start_time=task.metrics.arrival_time, ), ) is_prefill = task.disaggregate_info is not None and task.disaggregate_info["role"] == "prefill" From 004c57fc570574dcc1781ed6fd436c1ee10defb1 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Tue, 13 Jan 2026 11:13:47 +0800 Subject: [PATCH 05/25] support engine <==> worker communication(request&response) --- fastdeploy/engine/common_engine.py | 67 ++++++++++++++----- .../engine/sched/resource_manager_v1.py | 9 ++- fastdeploy/inter_communicator/fmq.py | 2 +- fastdeploy/scheduler/local_scheduler.py | 4 ++ fastdeploy/worker/gpu_model_runner.py | 3 + fastdeploy/worker/gpu_worker.py | 5 +- fastdeploy/worker/worker_process.py | 39 +++++++++++ 7 files changed, 109 insertions(+), 20 deletions(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index 16a8a219202..87dfbdc7a58 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -28,6 +28,7 @@ import time import traceback import weakref +import asyncio from concurrent.futures import ThreadPoolExecutor from typing import Dict, List, Optional, Tuple @@ -50,6 +51,7 @@ ZmqIpcServer, ZmqTcpServer, ) +from fastdeploy.inter_communicator.fmq import FMQ from fastdeploy.metrics.metrics import main_process_metrics from fastdeploy.model_executor.guided_decoding import schema_checker from fastdeploy.plugins.token_processor import load_token_processor_plugins @@ -82,9 +84,6 @@ def __init__(self, cfg, start_queue=True, use_async_llm=False): self.cfg = cfg self.use_async_llm = use_async_llm - self.is_paused = False # pause request generation - self._pause_cond = threading.Condition() - if self.cfg.parallel_config.data_parallel_size > 1: self.llm_logger = get_logger( "fastdeploy", f"fastdeploy_dprank{self.cfg.parallel_config.local_data_parallel_id}.log" @@ -92,6 +91,17 @@ def __init__(self, cfg, start_queue=True, use_async_llm=False): else: self.llm_logger = llm_logger + self.is_paused = False # pause request generation + self._pause_cond = threading.Condition() + + self._ctrl_worker_output_queues = [] + tp_size = cfg.parallel_config.tensor_parallel_size + dp_index = cfg.parallel_config.local_data_parallel_id + for rank in range(tp_size): + name = f"ctrl_w2e_rank{rank+tp_size*dp_index}" + self.llm_logger.info(f"Init Worker Control Output Queue: {name}(consumer)") + self._ctrl_worker_output_queues.append(FMQ().queue(name, "consumer")) + self.scheduler = cfg.scheduler_config.scheduler() self.enable_decode_cache_task = envs.FD_ENABLE_CACHE_TASK == "1" @@ -1072,7 +1082,7 @@ def _insert_zmq_task_to_scheduler(self): break if ControlRequest.is_control_request(data): - try: + try: #todo: run control request async, do not block request generation control_req = ControlRequest.from_dict(data) self.run_control_method(control_req) except Exception as e: @@ -1151,7 +1161,6 @@ def run_control_method(self, control_req: ControlRequest): - If no handler exists, returns error with available methods """ method = control_req.get_method() - args = control_req.get_args() request_id = control_req.request_id try: @@ -1165,7 +1174,7 @@ def run_control_method(self, control_req: ControlRequest): self.send_response_server.send_response(request_id, [error_result]) return - result = handler(args) + result = handler(control_req) self.llm_logger.info(f"SUCCESS run control method {method}.") succ_result = ControlResponse(request_id, 200, "Success", result) self.send_response_server.send_response(request_id, [succ_result]) @@ -1176,7 +1185,7 @@ def run_control_method(self, control_req: ControlRequest): error_result = ControlResponse(request_id, 500, error_msg) self.send_response_server.send_response(request_id, [error_result]) - def _control_pause(self, args: dict) -> dict | None: + def _control_pause(self, control_request: ControlRequest) -> dict | None: """暂停请求生成 Args: @@ -1187,6 +1196,11 @@ def _control_pause(self, args: dict) -> dict | None: - error_code: 错误代码,0表示成功,非0表示失败 - error_msg: 错误信息,成功时为空字符串 """ + if not envs.ENABLE_V1_KVCACHE_SCHEDULER: + raise Exception(f"pause only supported in ENABLE_V1_KVCACHE_SCHEDULER") + if self.cfg.scheduler_config.name != "local": + raise Exception(f"pause only supported in local scheduler, current {self.cfg.scheduler_config.name}") + with self._pause_cond: if self.is_paused: self.llm_logger.info("Pause Request Generation: already paused.") @@ -1212,14 +1226,12 @@ def _control_pause(self, args: dict) -> dict | None: self.resource_manager.get_real_bsz() self.engine_worker_queue.put_tasks((running_reqs, self.resource_manager.real_bsz)) self.resource_manager.wait_worker_inflight_requests_finish(timeout=60) - self.resource_manager.log_status() - self.engine_worker_queue.clear_data() + #self.engine_worker_queue.clear_data() self.token_processor.clear_data() - #self.resource_manager.clear_data() self.resource_manager.log_status() # abort inflight requests to user - inflight_requests = [req.raw for req in self.scheduler.requests.values()] + inflight_requests = self.scheduler.get_inflight_requests() self.llm_logger.info(f"Start Abort Inflight Requests, total {len(inflight_requests)} waiting requests") for req in inflight_requests: self._send_error_response(req.request_id, "Request is aborted since LLM Engine is paused.") @@ -1228,7 +1240,7 @@ def _control_pause(self, args: dict) -> dict | None: self.resource_manager.cache_manager.reset() return None - def _control_resume(self, args: dict) -> dict | None: + def _control_resume(self, control_request: ControlRequest) -> dict | None: """恢复暂停的请求生成 Args: @@ -1246,7 +1258,7 @@ def _control_resume(self, args: dict) -> dict | None: self.llm_logger.info(f"END Resume Request Generation") return None - def _control_is_paused(self, args: dict) -> bool: + def _control_is_paused(self, control_request: ControlRequest) -> bool: """检查是否暂停了请求生成 Args: @@ -1259,7 +1271,7 @@ def _control_is_paused(self, args: dict) -> bool: with self._pause_cond: return {"is_paused": self.is_paused} - def _control_update_weights(self, args: dict) -> dict | None: + def _control_update_weights(self, control_request: ControlRequest) -> dict | None: """更新模型权重 Args: @@ -1269,7 +1281,32 @@ def _control_update_weights(self, args: dict) -> dict | None: dict | None: 返回结果字典或None,包含更新权重的操作结果信息 """ self.llm_logger.info(f"Update Model Weights") - return None + with self._pause_cond: + if self.is_paused is False: + error_msg = f"Pause LLM Engine first before calling updating weights" + self.llm_logger.error(error_msg) + raise Exception(error_msg) + return self._call_worker(control_request, 60) + + def _call_worker(self, control_request: ControlRequest, timeout: int): + request_id = control_request.request_id + self.engine_worker_queue.put_tasks(([control_request], 1)) + + responses = [] + for output_queue in self._ctrl_worker_output_queues: + msg = asyncio.run(output_queue.get(timeout=timeout*1000)) # todo: fix timeout when tp > 1 + if msg is None: + raise Exception("Worker Update Weights Timeouted after 600s") + response: ControlResponse = msg.payload + if response.request_id != request_id: + self.llm_logger.info(f"ignore old control response from worker:{output_queue.name} {response}") + continue + if response.error_code != 200: + self.llm_logger.info(f"Call Worker Failed: {output_queue.name} {response.error_message}") + raise Exception(f"Call Worker error: {response.error_message}") + self.llm_logger.info(f"Call Worker Succeed: {output_queue.name} {response.result}") + responses.append(response.result) + return responses def _send_error_response(self, request_id, error_msg, error_code: int = 500): self.llm_logger.error( diff --git a/fastdeploy/engine/sched/resource_manager_v1.py b/fastdeploy/engine/sched/resource_manager_v1.py index f102a9a5890..1cd570cb0b1 100644 --- a/fastdeploy/engine/sched/resource_manager_v1.py +++ b/fastdeploy/engine/sched/resource_manager_v1.py @@ -258,10 +258,12 @@ def preempted_all(self): preempted_reqs = [] for i in range(len(self.running)): req = self.running.pop() + # txt2image: req.use_extend_tables is True, req can not be preempted. txt2image is not used in RL. + if req.use_extend_tables: + self.running.insert(0, req) + continue req.status = RequestStatus.PREEMPTED req.num_computed_tokens = 0 - #self.tasks_list[req.idx] = None - #self.stop_flags[req.idx] = True self._free_blocks(req) req.cached_block_num = 0 self.to_be_rescheduled_request_id_set.add(req.request_id) @@ -271,7 +273,8 @@ def preempted_all(self): def wait_worker_inflight_requests_finish(self, timeout=60): count = 0 while count < timeout * 1000: - running_reqs_count = len(self.to_be_rescheduled_request_id_set) + # wait ongoing running and rescheduled requests finished in worker + running_reqs_count = len(self.to_be_rescheduled_request_id_set) + len(self.running) if running_reqs_count == 0: break diff --git a/fastdeploy/inter_communicator/fmq.py b/fastdeploy/inter_communicator/fmq.py index f2c98196c99..e30e4f54e6d 100644 --- a/fastdeploy/inter_communicator/fmq.py +++ b/fastdeploy/inter_communicator/fmq.py @@ -214,7 +214,7 @@ def __init__(self, context, name: str, role: str = "producer"): else: self.socket.bind(full_ep) - fmq_logger.info(f"Queue {name} initialized on {full_ep}") + fmq_logger.info(f"Queue {name}({role}) initialized on {full_ep}") async def put(self, data: Any, shm_threshold: int = 1024 * 1024): """ diff --git a/fastdeploy/scheduler/local_scheduler.py b/fastdeploy/scheduler/local_scheduler.py index 9bfe90a6e29..fc4a64686b5 100644 --- a/fastdeploy/scheduler/local_scheduler.py +++ b/fastdeploy/scheduler/local_scheduler.py @@ -158,6 +158,10 @@ def _recycle(self, request_id: Optional[str] = None): else: self.ids_read_cursor -= len(expired_ids) + def get_inflight_requests(self) -> List[Request]: + with self.mutex: + return [request.raw for request in self.requests.values()] + def put_requests(self, requests: List[Request]) -> List[Tuple[str, Optional[str]]]: """ Add new requests to the scheduler queue. diff --git a/fastdeploy/worker/gpu_model_runner.py b/fastdeploy/worker/gpu_model_runner.py index 41eac4b6b4f..60bd93a6d06 100644 --- a/fastdeploy/worker/gpu_model_runner.py +++ b/fastdeploy/worker/gpu_model_runner.py @@ -2810,6 +2810,9 @@ def update_parameters(self, pid): self.dynamic_weight_manager._log_memory("dynamic weight manager update all memory") + def update_weights(self): + logger.info("GPU Model Runner update weights inplace") + def padding_cudagraph_inputs(self) -> None: """ Clean buffers used for the CUDA graph when replaying the CUDA graph with the padded batch. diff --git a/fastdeploy/worker/gpu_worker.py b/fastdeploy/worker/gpu_worker.py index 0d57ccf2504..2932b7cc92b 100644 --- a/fastdeploy/worker/gpu_worker.py +++ b/fastdeploy/worker/gpu_worker.py @@ -184,6 +184,9 @@ def initialize_cache(self, num_gpu_blocks: int) -> None: # accurate cache size self.model_runner.update_share_input_block_num(num_gpu_blocks=num_gpu_blocks) + def update_weights(self): + self.model_runner.update_weights() + def execute_model( self, model_forward_batch: Optional[List[Request]] = None, @@ -220,4 +223,4 @@ def check_health(self) -> bool: def cal_theortical_kvcache(self) -> int: """Calculate the block memory required""" - return self.model_runner.cal_theortical_kvcache() + return self.model_runner.cal_theortical_kvcache() \ No newline at end of file diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 051e5d59338..14d7883f345 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -18,7 +18,9 @@ import json import os import time +import traceback from typing import Tuple +import asyncio import numpy as np import paddle @@ -56,12 +58,14 @@ ModelWeightsStatus, RearrangeExpertStatus, ) +from fastdeploy.inter_communicator.fmq import FMQ from fastdeploy.model_executor.layers.quantization import parse_quant_config from fastdeploy.model_executor.utils import v1_loader_support from fastdeploy.platforms import current_platform from fastdeploy.scheduler import SchedulerConfig from fastdeploy.utils import get_logger, optional_type from fastdeploy.worker.worker_base import WorkerBase +from fastdeploy.engine.request import ControlRequest, ControlResponse logger = get_logger("worker_process", "worker_process.log") @@ -163,6 +167,11 @@ def __init__(self, fd_config: FDConfig, ranks: int = 1, local_rank: int = 0) -> self.max_chips_per_node = 16 if current_platform.is_iluvatar() else 8 + def init_control(self): + queue_name = f"ctrl_w2e_rank{self.local_rank}" + logger.info(f"Init Control Output Queue: {queue_name}(producer)") + self._ctrl_output = FMQ().queue(queue_name, "producer") + def init_health_status(self) -> None: """ Initialize the health status of the worker. @@ -488,6 +497,11 @@ def event_loop_normal(self) -> None: num_running_requests = int(bsz) req_dicts.extend(req_dict) + # try run control method + if len(req_dicts) == 1 and isinstance(req_dicts[0], ControlRequest): + self.run_control_method(req_dicts[0]) + continue + req_ids = [req.request_id for req in req_dicts] logger.info( f"Rank: {self.local_rank}, num_running_requests: {num_running_requests}, " @@ -618,6 +632,30 @@ def load_model(self) -> None: paddle.distributed.barrier() self.loaded_model_signal.value[0] = 1 + def run_control_method(self, control_request: ControlRequest) -> None: + request_id = control_request.request_id + method = control_request.method + kwargs = control_request.args + + handler = getattr(self.worker, method, None) + if handler is None or not callable(handler): + error_result = ControlResponse(request_id, 400, f"unknown control method {method}") + self._send_control_response(error_result) + return + + try: + result = handler(**kwargs) + succ_result = ControlResponse(request_id, 200, "Success", result) + self._send_control_response(succ_result) + except Exception as e: + error_msg = f"Failed run control method {method}: {str(e)}" + logger.info(f"{error_msg}\n{traceback.format_exc()}") + error_result = ControlResponse(request_id, 500, error_msg) + self._send_control_response(error_result) + + def _send_control_response(self, control_response: ControlResponse): + logger.info(f"Rank-{self.local_rank} put control output {control_response} to engine") + asyncio.run(self._ctrl_output.put(control_response, shm_threshold=100*1024*1024)) def parse_args(): """ @@ -1058,6 +1096,7 @@ def run_worker_proc() -> None: worker_proc = IluvatarPaddleDisWorkerProc(fd_config, ranks, local_rank) else: worker_proc = PaddleDisWorkerProc(fd_config, ranks, local_rank) + worker_proc.init_control() # Initialize device and create model runner worker_proc.init_device() From 3f1115cc9bc42a20e008cf0b2a6c43045ce342d9 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Wed, 14 Jan 2026 20:46:23 +0800 Subject: [PATCH 06/25] support sync weights through RDMA from checkpoint_transfer --- fastdeploy/rl/dynamic_weight_manager.py | 68 ++++++++++++++++++++++++- fastdeploy/worker/gpu_model_runner.py | 4 +- fastdeploy/worker/worker_process.py | 2 +- 3 files changed, 69 insertions(+), 5 deletions(-) diff --git a/fastdeploy/rl/dynamic_weight_manager.py b/fastdeploy/rl/dynamic_weight_manager.py index ee9dbb892d2..6942219c647 100644 --- a/fastdeploy/rl/dynamic_weight_manager.py +++ b/fastdeploy/rl/dynamic_weight_manager.py @@ -26,14 +26,40 @@ from fastdeploy.config import FDConfig from fastdeploy.inter_communicator import ModelWeightsStatus +def sync_weights_by_rdma(step, rank): + etcd_server = "127.0.0.1:2379" + + from checkpoint_transfer.core import RDMAWeightsDownloader + import io + config = { "etcd_server": etcd_server } + downloader = RDMAWeightsDownloader(config) + downloader.initialize() + logger.info(f"Fetching weights for step:{step}, rank:{rank}...") + data = downloader.get_weights(step, rank) + if data is None: + logger.error("Failed to get weights!") + logger.info(f"Successfully retrieved data. Type: {type(data)}") + if isinstance(data, np.ndarray): + data_bytes = data.tobytes() + elif isinstance(data, (bytes, bytearray)): + data_bytes = data + else: + data_bytes = bytes(data) + logger.info(f"Data size: {len(data_bytes)} bytes") + + buffer = io.BytesIO(data_bytes) + new_state_dict = paddle.load(buffer) + return new_state_dict + class DynamicWeightManager: """Manages model weights loading, updating and shared state across processes.""" - def __init__(self, fd_config: FDConfig, models): + def __init__(self, fd_config: FDConfig, models, local_rank: int): """Initialize with config and model instances.""" self.fd_config = fd_config self.load_config = fd_config.load_config + self.local_rank = local_rank self.parallel_config = fd_config.parallel_config self.state_dict: Dict[str, paddle.Tensor] = {} self.rank = fd_config.parallel_config.tensor_parallel_rank @@ -46,7 +72,11 @@ def __init__(self, fd_config: FDConfig, models): else: self.model_list = models self._capture_model_state() - self.update_parameters() + if self.load_config.load_strategy == "rsync": + step = 100 # todo read from {model}/version.txt + self.update_weights_by_rdma(step, self.local_rank) + else: + self.update_parameters() self.finalize_update() logger.info( @@ -62,6 +92,40 @@ def _capture_model_state(self): logger.info(f"Model param: {name}, shape={param.shape}, dtype={param.dtype}") self.state_dict[name] = param + def update_weights_by_rdma(self, step, rank): + old_state_dict = self.state_dict + def valid_parameters(old_state_dict, new_state_dict): + is_valid = True + for key in old_state_dict: + if key not in new_state_dict: + is_valid = False + logger.error(f"Invalid parameter: {key} not in new_state_dict") + elif old_state_dict[key].shape != new_state_dict[key].shape: + is_valid = False + logger.error(f"Invalid parameter: {key} shape mismatch, " + f"new shape:{new_state_dict[key].shape}, " + f"old shape:{old_state_dict[key].shape}") + elif old_state_dict[key].dtype != new_state_dict[key].dtype: + is_valid = False + logger.error(f"Invalid parameter: {key} dtype mismatch") + return is_valid + + start_time = time.perf_counter() + + new_state_dict = sync_weights_by_rdma(step, rank) + #old_state_dict = self.model_list[9].state_dict() + if not valid_parameters(old_state_dict, new_state_dict): + logger.error("Invalid new_state_dict, update parameters failed") + return + + assign_start = time.perf_counter() + for name, param in old_state_dict.items(): + param.set_value(new_state_dict[name]) + logger.info(f"params set value cost {time.perf_counter()-assign_start:.2f} seconds") + + logger.info(f"update weights inplace cost {time.perf_counter()-start_time:.2f} seconds") + + def update_parameters(self, pid: int = 0, restart_process_group=False) -> None: """Core method to update model parameters based on strategy.""" start_time = time.perf_counter() diff --git a/fastdeploy/worker/gpu_model_runner.py b/fastdeploy/worker/gpu_model_runner.py index 60bd93a6d06..e41c161701e 100644 --- a/fastdeploy/worker/gpu_model_runner.py +++ b/fastdeploy/worker/gpu_model_runner.py @@ -1503,8 +1503,7 @@ def load_model(self) -> None: # 1.1 Load RL dynamic model if self.fd_config.load_config.dynamic_load_weight: from fastdeploy.rl.dynamic_weight_manager import DynamicWeightManager - - self.dynamic_weight_manager = DynamicWeightManager(self.fd_config, self.model) + self.dynamic_weight_manager = DynamicWeightManager(self.fd_config, self.model, self.local_rank) # 2. Load lora model @@ -2812,6 +2811,7 @@ def update_parameters(self, pid): def update_weights(self): logger.info("GPU Model Runner update weights inplace") + self.dynamic_weight_manager.update_weights_by_rdma(100, self.local_rank) def padding_cudagraph_inputs(self) -> None: """ diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 14d7883f345..252ca122a40 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -814,7 +814,7 @@ def parse_args(): parser.add_argument( "--load_strategy", type=str, - choices=["ipc", "ipc_snapshot", "meta", "normal"], + choices=["ipc", "ipc_snapshot", "meta", "normal", "rsync"], default="ipc_snapshot", help="Weight loading method when dynamic loading is enabled: " "'ipc': real-time IPC streaming with automatic resharding, " From ae306bcc557ea8404917d8ad3efc2920d73e3f38 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 15 Jan 2026 17:33:09 +0800 Subject: [PATCH 07/25] support specified version, rsync_config in update_weights rpc call --- fastdeploy/config.py | 3 +- fastdeploy/engine/args_utils.py | 10 ++++ fastdeploy/engine/common_engine.py | 1 + fastdeploy/engine/engine.py | 1 + fastdeploy/entrypoints/openai/api_server.py | 16 ++++++- fastdeploy/rl/dynamic_weight_manager.py | 53 +++++++++++++++------ fastdeploy/worker/gpu_model_runner.py | 7 ++- fastdeploy/worker/gpu_worker.py | 7 +-- fastdeploy/worker/worker_process.py | 26 +++++++--- 9 files changed, 93 insertions(+), 31 deletions(-) diff --git a/fastdeploy/config.py b/fastdeploy/config.py index d5058d6b3ca..6a58874dd9a 100644 --- a/fastdeploy/config.py +++ b/fastdeploy/config.py @@ -1173,6 +1173,7 @@ def __init__( self.load_choices: Union[str, LoadChoices] = LoadChoices.DEFAULT.value self.dynamic_load_weight: bool = False self.load_strategy: Optional[Literal["ipc", "ipc_snapshot", "meta", "normal"]] = "normal" + self.rsync_config: Optional[Dict[str, Any]] = None for key, value in args.items(): if hasattr(self, key): setattr(self, key, value) @@ -1267,7 +1268,7 @@ def print(self): """ Print all configuration information. """ - logger.info("EPLB Configuration Information :") + i.info("EPLB Configuration Information :") for k, v in self.__dict__.items(): logger.info("{:<20}:{:<6}{}".format(k, "", v)) logger.info("=============================================================") diff --git a/fastdeploy/engine/args_utils.py b/fastdeploy/engine/args_utils.py index b7d7f4d807e..80762f9b662 100644 --- a/fastdeploy/engine/args_utils.py +++ b/fastdeploy/engine/args_utils.py @@ -184,6 +184,10 @@ class EngineArgs: """ dynamic load weight strategy """ + rsync_config: Optional[Dict[str, Any]] = None + """ + rsync weights config info + """ quantization: Optional[Dict[str, Any]] = None guided_decoding_backend: str = "off" """ @@ -793,6 +797,12 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: default=EngineArgs.load_strategy, help="Flag to dynamic load strategy.", ) + model_group.add_argument( + "--rsync-config", + type=json.loads, + default=EngineArgs.rsync_config, + help="Rsync weights config", + ) model_group.add_argument( "--engine-worker-queue-port", type=lambda s: s.split(",") if s else None, diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index 87dfbdc7a58..5bedb1d724e 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -1871,6 +1871,7 @@ def _start_worker_service(self): f" --graph_optimization_config '{self.cfg.graph_opt_config.to_json_string()}'" f" --guided_decoding_backend {self.cfg.structured_outputs_config.guided_decoding_backend}" f" --load_strategy {self.cfg.load_config.load_strategy}" + f" --rsync_config '{json.dumps(self.cfg.load_config.rsync_config)}'" f" --early_stop_config '{self.cfg.early_stop_config.to_json_string()}'" f" --reasoning_parser {self.cfg.structured_outputs_config.reasoning_parser}" f" --load_choices {self.cfg.load_config.load_choices}" diff --git a/fastdeploy/engine/engine.py b/fastdeploy/engine/engine.py index 48f9e5a1c25..cb9440c06f3 100644 --- a/fastdeploy/engine/engine.py +++ b/fastdeploy/engine/engine.py @@ -559,6 +559,7 @@ def _start_worker_service(self): f" --graph_optimization_config '{self.cfg.graph_opt_config.to_json_string()}'" f" --guided_decoding_backend {self.cfg.structured_outputs_config.guided_decoding_backend}" f" --load_strategy {self.cfg.load_config.load_strategy}" + f" --rsync_config '{json.dumps(self.cfg.load_config.rsync_config)}'" f" --early_stop_config '{self.cfg.early_stop_config.to_json_string()}'" f" --reasoning_parser {self.cfg.structured_outputs_config.reasoning_parser}" f" --load_choices {self.cfg.load_config.load_choices}" diff --git a/fastdeploy/entrypoints/openai/api_server.py b/fastdeploy/entrypoints/openai/api_server.py index 29d6c04b407..cb5faf51242 100644 --- a/fastdeploy/entrypoints/openai/api_server.py +++ b/fastdeploy/entrypoints/openai/api_server.py @@ -392,7 +392,21 @@ async def is_paused(request: Request) -> Response: @app.post("/v1/update_weights") async def update_weights(request: Request) -> Response: request_id = f"control-{uuid.uuid4()}" - control_request = ControlRequest(request_id, "update_weights") + # 从请求体中获取参数 + request_data = await request.json() + + # 提取version和rsync_config参数 + version = request_data.get("version") + rsync_config = request_data.get("rsync_config") + + # 构建控制请求,传递参数 + args = {} + if version is not None: + args["version"] = version + if rsync_config is not None: + args["rsync_config"] = rsync_config + + control_request = ControlRequest(request_id, "update_weights", args) control_response = await app.state.engine_client.run_control_method(control_request) return control_response.to_api_json_response() diff --git a/fastdeploy/rl/dynamic_weight_manager.py b/fastdeploy/rl/dynamic_weight_manager.py index 6942219c647..c44fd6ec3d6 100644 --- a/fastdeploy/rl/dynamic_weight_manager.py +++ b/fastdeploy/rl/dynamic_weight_manager.py @@ -15,6 +15,7 @@ """ import os +import io import time from multiprocessing.shared_memory import SharedMemory from typing import Any, Dict, List @@ -26,18 +27,15 @@ from fastdeploy.config import FDConfig from fastdeploy.inter_communicator import ModelWeightsStatus -def sync_weights_by_rdma(step, rank): - etcd_server = "127.0.0.1:2379" - +def sync_weights_by_rdma(config, step, rank): from checkpoint_transfer.core import RDMAWeightsDownloader - import io - config = { "etcd_server": etcd_server } downloader = RDMAWeightsDownloader(config) downloader.initialize() logger.info(f"Fetching weights for step:{step}, rank:{rank}...") data = downloader.get_weights(step, rank) if data is None: logger.error("Failed to get weights!") + raise Exception("Failed to rsync weights through checkpoint_transfer") logger.info(f"Successfully retrieved data. Type: {type(data)}") if isinstance(data, np.ndarray): data_bytes = data.tobytes() @@ -73,8 +71,7 @@ def __init__(self, fd_config: FDConfig, models, local_rank: int): self.model_list = models self._capture_model_state() if self.load_config.load_strategy == "rsync": - step = 100 # todo read from {model}/version.txt - self.update_weights_by_rdma(step, self.local_rank) + self.update_weights_by_rdma() else: self.update_parameters() self.finalize_update() @@ -92,8 +89,7 @@ def _capture_model_state(self): logger.info(f"Model param: {name}, shape={param.shape}, dtype={param.dtype}") self.state_dict[name] = param - def update_weights_by_rdma(self, step, rank): - old_state_dict = self.state_dict + def update_weights_by_rdma(self, version: str = None, rsync_config: dict[str, Any] = None): def valid_parameters(old_state_dict, new_state_dict): is_valid = True for key in old_state_dict: @@ -110,20 +106,41 @@ def valid_parameters(old_state_dict, new_state_dict): logger.error(f"Invalid parameter: {key} dtype mismatch") return is_valid - start_time = time.perf_counter() + if rsync_config is None: + rsync_config = self.fd_config.load_config.rsync_config + if rsync_config is None or len(rsync_config) == "": + raise Exception(f"rsync config not set, please set it in 1) launch arguments '--rsync-config' " + f"or 2) interface arguments 'rsync_config'") + + if version is None or version == "": + version = self.read_model_version_from_file() + if version is None or version == "": + raise Exception(f"rsync model version not set, please set it in 1) {{model_version}}/version.txt " + f"or 2) interface arguments 'version'") - new_state_dict = sync_weights_by_rdma(step, rank) - #old_state_dict = self.model_list[9].state_dict() + llm_logger.info(f"START update_weights_by_rdma, version:{version}, rsync_config:{rsync_config}") + rank = self.local_rank + + sync_start = time.perf_counter() + new_state_dict = sync_weights_by_rdma(rsync_config, version, rank) + sync_cost = time.perf_counter() - sync_start + logger.info(f"weights sync cost {sync_cost:.2f} seconds") + + old_state_dict = self.state_dict if not valid_parameters(old_state_dict, new_state_dict): logger.error("Invalid new_state_dict, update parameters failed") return - assign_start = time.perf_counter() + update_start = time.perf_counter() for name, param in old_state_dict.items(): param.set_value(new_state_dict[name]) - logger.info(f"params set value cost {time.perf_counter()-assign_start:.2f} seconds") + update_cost = time.perf_counter() - update_start + logger.info(f"params set value cost {update_cost:.2f} seconds") - logger.info(f"update weights inplace cost {time.perf_counter()-start_time:.2f} seconds") + total_cost = time.perf_counter() - sync_start + logger.info(f"END update_weights_by_rdma, cost {total_cost:.2f} seconds", + f" version:{version}, rsync_config: {rsync_config}") + return {"sync_cost": sync_cost, "update_cost": update_cost, "total_cost": total_cost, "version": version, "rank": rank} def update_parameters(self, pid: int = 0, restart_process_group=False) -> None: @@ -321,6 +338,12 @@ def _update_shared_status(self, pid: int, status: int) -> None: if self.rank == 0: value[self.rank] = status + def read_model_version_from_file(self): + model_dir = self.fd_config.model_config.model + with open(os.path.join(model_dir, "version.txt")) as f: + version = f.read().strip() + return version + @staticmethod def check_model_weights_status(model_weights_status, model_runner, pid, block): """ diff --git a/fastdeploy/worker/gpu_model_runner.py b/fastdeploy/worker/gpu_model_runner.py index e41c161701e..646fe260ab1 100644 --- a/fastdeploy/worker/gpu_model_runner.py +++ b/fastdeploy/worker/gpu_model_runner.py @@ -20,7 +20,7 @@ import time from concurrent.futures import Future from threading import Thread -from typing import List, Optional, cast +from typing import Any, List, Optional, cast, Dict import numpy as np import paddle @@ -2809,9 +2809,8 @@ def update_parameters(self, pid): self.dynamic_weight_manager._log_memory("dynamic weight manager update all memory") - def update_weights(self): - logger.info("GPU Model Runner update weights inplace") - self.dynamic_weight_manager.update_weights_by_rdma(100, self.local_rank) + def update_weights(self, version: str = None, rsync_config: Dict[str, Any] = None): + return self.dynamic_weight_manager.update_weights_by_rdma(version, rsync_config) def padding_cudagraph_inputs(self) -> None: """ diff --git a/fastdeploy/worker/gpu_worker.py b/fastdeploy/worker/gpu_worker.py index 2932b7cc92b..c98ad048593 100644 --- a/fastdeploy/worker/gpu_worker.py +++ b/fastdeploy/worker/gpu_worker.py @@ -16,7 +16,7 @@ import gc import time -from typing import List, Optional +from typing import List, Optional, Any, Dict import paddle import pynvml @@ -184,8 +184,9 @@ def initialize_cache(self, num_gpu_blocks: int) -> None: # accurate cache size self.model_runner.update_share_input_block_num(num_gpu_blocks=num_gpu_blocks) - def update_weights(self): - self.model_runner.update_weights() + def update_weights(self, version: str = None, rsync_config: Dict[str, Any] = None): + """update weights in place""" + return self.model_runner.update_weights(version, rsync_config) def execute_model( self, diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 252ca122a40..f31a7053488 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -492,15 +492,18 @@ def event_loop_normal(self) -> None: else: self.exist_task_signal.value[0] = ExistTaskStatus.EMPTY - req_dicts = [] + req_dicts, control_reqs = [], [] for req_dict, bsz in tasks: - num_running_requests = int(bsz) - req_dicts.extend(req_dict) + if len(req_dict) > 0 and isinstance(req_dict[0], ControlRequest): + control_reqs.append(req_dict[0]) + else: + num_running_requests = int(bsz) + req_dicts.extend(req_dict) - # try run control method - if len(req_dicts) == 1 and isinstance(req_dicts[0], ControlRequest): - self.run_control_method(req_dicts[0]) - continue + if len(control_reqs) > 0: + logger.info(f"Rank: {self.local_rank} received {len(control_reqs)} control request.") + for control_req in control_reqs: + self.run_control_method(control_req) req_ids = [req.request_id for req in req_dicts] logger.info( @@ -633,6 +636,7 @@ def load_model(self) -> None: self.loaded_model_signal.value[0] = 1 def run_control_method(self, control_request: ControlRequest) -> None: + logger.info(f"Start run control request: {control_request}") request_id = control_request.request_id method = control_request.method kwargs = control_request.args @@ -647,6 +651,7 @@ def run_control_method(self, control_request: ControlRequest) -> None: result = handler(**kwargs) succ_result = ControlResponse(request_id, 200, "Success", result) self._send_control_response(succ_result) + logger.info(f"Success run control request: {control_request}, response: {succ_result}") except Exception as e: error_msg = f"Failed run control method {method}: {str(e)}" logger.info(f"{error_msg}\n{traceback.format_exc()}") @@ -820,6 +825,12 @@ def parse_args(): "'ipc': real-time IPC streaming with automatic resharding, " "'ipc_snapshot': load from disk snapshot of IPC weights.", ) + parser.add_argument( + "--rsync_config", + type=json.loads, + default=None, + help="Rsync weights config", + ) parser.add_argument( "--enable_logprob", action="store_true", @@ -1029,6 +1040,7 @@ def initialize_fd_config(args, ranks: int = 1, local_rank: int = 0) -> FDConfig: logger.info(f"- Dynamic load weight: {load_config.dynamic_load_weight}") logger.info(f"- Load strategy: {load_config.load_strategy}") + logger.info(f"- Rsync config: {load_config.rsync_config}, {type(load_config.rsync_config)}") if not ( current_platform.is_cuda() From 1b1a3484a23cb8ed1f5a3e0f2b51bd1044fe9b3b Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 15 Jan 2026 17:51:50 +0800 Subject: [PATCH 08/25] add pause, update_weights, resume interface for async RL --- fastdeploy/engine/common_engine.py | 48 ++++++++++-------- fastdeploy/engine/request.py | 49 +++++++++---------- .../engine/sched/resource_manager_v1.py | 28 ++++++----- fastdeploy/entrypoints/engine_client.py | 5 +- fastdeploy/entrypoints/openai/api_server.py | 15 ++++-- fastdeploy/rl/dynamic_weight_manager.py | 43 ++++++++++------ fastdeploy/worker/gpu_model_runner.py | 3 +- fastdeploy/worker/gpu_worker.py | 4 +- fastdeploy/worker/worker_process.py | 7 +-- 9 files changed, 115 insertions(+), 87 deletions(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index 5bedb1d724e..59813e1df7a 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -16,6 +16,7 @@ from __future__ import annotations +import asyncio import copy import json import multiprocessing @@ -28,7 +29,6 @@ import time import traceback import weakref -import asyncio from concurrent.futures import ThreadPoolExecutor from typing import Dict, List, Optional, Tuple @@ -39,7 +39,13 @@ from tqdm import tqdm import fastdeploy.metrics.trace as tracing -from fastdeploy.engine.request import Request, ControlRequest, ControlResponse, RequestOutput, RequestType +from fastdeploy.engine.request import ( + ControlRequest, + ControlResponse, + Request, + RequestOutput, + RequestType, +) from fastdeploy.engine.resource_manager import ResourceManager from fastdeploy.engine.sched.resource_manager_v1 import ResourceManagerV1 from fastdeploy.eplb.utils import init_eplb_signals @@ -83,7 +89,7 @@ def __init__(self, cfg, start_queue=True, use_async_llm=False): """ self.cfg = cfg self.use_async_llm = use_async_llm - + if self.cfg.parallel_config.data_parallel_size > 1: self.llm_logger = get_logger( "fastdeploy", f"fastdeploy_dprank{self.cfg.parallel_config.local_data_parallel_id}.log" @@ -91,7 +97,7 @@ def __init__(self, cfg, start_queue=True, use_async_llm=False): else: self.llm_logger = llm_logger - self.is_paused = False # pause request generation + self.is_paused = False # pause request generation self._pause_cond = threading.Condition() self._ctrl_worker_output_queues = [] @@ -1082,7 +1088,7 @@ def _insert_zmq_task_to_scheduler(self): break if ControlRequest.is_control_request(data): - try: #todo: run control request async, do not block request generation + try: # todo: run control request async, do not block request generation control_req = ControlRequest.from_dict(data) self.run_control_method(control_req) except Exception as e: @@ -1107,7 +1113,9 @@ def _insert_zmq_task_to_scheduler(self): if self.is_paused: self.llm_logger.warning(f"Engine is paused, drop request: {request}") - self._send_error_response(request.request_id, "Request is aborted since LLM Engine is paused.") + self._send_error_response( + request.request_id, "Request is aborted since LLM Engine is paused." + ) continue except Exception as e: self.llm_logger.error(f"Receive request error: {e}, {traceback.format_exc()!s}") @@ -1151,10 +1159,10 @@ def _insert_zmq_task_to_scheduler(self): def run_control_method(self, control_req: ControlRequest): """ Execute control methods for engine management using dynamic method invocation. - + Args: control_req: ControlRequest instance containing method name and arguments - + Usage: - Control request with method "get_metrics" will call self._control_get_metrics(args) - Method names are automatically mapped to handler methods with prefix "_control_" @@ -1162,10 +1170,10 @@ def run_control_method(self, control_req: ControlRequest): """ method = control_req.get_method() request_id = control_req.request_id - + try: self.llm_logger.info(f"START run control method {request_id}: {method}") - + handler_name = f"_control_{method}" handler = getattr(self, handler_name, None) if handler is None or not callable(handler): @@ -1173,12 +1181,12 @@ def run_control_method(self, control_req: ControlRequest): self.llm_logger.error(str(error_result)) self.send_response_server.send_response(request_id, [error_result]) return - + result = handler(control_req) self.llm_logger.info(f"SUCCESS run control method {method}.") succ_result = ControlResponse(request_id, 200, "Success", result) self.send_response_server.send_response(request_id, [succ_result]) - + except Exception as e: error_msg = f"Failed run control method {method}: {str(e)}" self.llm_logger.error(f"{error_msg}\n{traceback.format_exc()}") @@ -1197,7 +1205,7 @@ def _control_pause(self, control_request: ControlRequest) -> dict | None: - error_msg: 错误信息,成功时为空字符串 """ if not envs.ENABLE_V1_KVCACHE_SCHEDULER: - raise Exception(f"pause only supported in ENABLE_V1_KVCACHE_SCHEDULER") + raise Exception("pause only supported in ENABLE_V1_KVCACHE_SCHEDULER") if self.cfg.scheduler_config.name != "local": raise Exception(f"pause only supported in local scheduler, current {self.cfg.scheduler_config.name}") @@ -1206,7 +1214,7 @@ def _control_pause(self, control_request: ControlRequest) -> dict | None: self.llm_logger.info("Pause Request Generation: already paused.") self.is_paused = True - self.llm_logger.info(f"Start Abort Running Requests") + self.llm_logger.info("Start Abort Running Requests") self.resource_manager.log_status() # preempted all running reqs. preempted reqs will be append to ResourceManager.waiting queue @@ -1226,7 +1234,7 @@ def _control_pause(self, control_request: ControlRequest) -> dict | None: self.resource_manager.get_real_bsz() self.engine_worker_queue.put_tasks((running_reqs, self.resource_manager.real_bsz)) self.resource_manager.wait_worker_inflight_requests_finish(timeout=60) - #self.engine_worker_queue.clear_data() + # self.engine_worker_queue.clear_data() self.token_processor.clear_data() self.resource_manager.log_status() @@ -1249,13 +1257,13 @@ def _control_resume(self, control_request: ControlRequest) -> dict | None: Returns: dict | None: 返回结果字典或None,包含恢复操作的状态信息 """ - self.llm_logger.info(f"START Resume Request Generation") + self.llm_logger.info("START Resume Request Generation") with self._pause_cond: if not self.is_paused: self.llm_logger.info("Resume Request Generation: not paused.") self.is_paused = False self._pause_cond.notify_all() - self.llm_logger.info(f"END Resume Request Generation") + self.llm_logger.info("END Resume Request Generation") return None def _control_is_paused(self, control_request: ControlRequest) -> bool: @@ -1280,10 +1288,10 @@ def _control_update_weights(self, control_request: ControlRequest) -> dict | Non Returns: dict | None: 返回结果字典或None,包含更新权重的操作结果信息 """ - self.llm_logger.info(f"Update Model Weights") + self.llm_logger.info("Update Model Weights") with self._pause_cond: if self.is_paused is False: - error_msg = f"Pause LLM Engine first before calling updating weights" + error_msg = "Pause LLM Engine first before calling updating weights" self.llm_logger.error(error_msg) raise Exception(error_msg) return self._call_worker(control_request, 60) @@ -1294,7 +1302,7 @@ def _call_worker(self, control_request: ControlRequest, timeout: int): responses = [] for output_queue in self._ctrl_worker_output_queues: - msg = asyncio.run(output_queue.get(timeout=timeout*1000)) # todo: fix timeout when tp > 1 + msg = asyncio.run(output_queue.get(timeout=timeout * 1000)) # todo: fix timeout when tp > 1 if msg is None: raise Exception("Worker Update Weights Timeouted after 600s") response: ControlResponse = msg.payload diff --git a/fastdeploy/engine/request.py b/fastdeploy/engine/request.py index 3d16dc17bfd..d567145c4cd 100644 --- a/fastdeploy/engine/request.py +++ b/fastdeploy/engine/request.py @@ -17,13 +17,12 @@ from __future__ import annotations import time -import traceback from dataclasses import asdict, dataclass, fields from enum import Enum from typing import Any, Dict, Generic, Optional, Union -from fastapi.responses import JSONResponse import numpy as np +from fastapi.responses import JSONResponse from typing_extensions import TypeVar from fastdeploy import envs @@ -380,9 +379,10 @@ def __repr__(self) -> str: except Exception as e: return f"" + class ControlRequest: """A generic control request that supports method and args for control operations. - + This request type is used for system-level control operations rather than typical inference requests. It enables dynamic control of engine behavior, resource management, and system configuration via a flexible method-args interface. @@ -407,19 +407,11 @@ def __init__( @classmethod def from_dict(cls, d: dict): """Create ControlRequest instance from dictionary.""" - return cls( - request_id=d["request_id"], - method=d["method"], - args=d.get("args", {}) - ) + return cls(request_id=d["request_id"], method=d["method"], args=d.get("args", {})) def to_dict(self) -> dict: """Convert ControlRequest into a serializable dict.""" - return { - "request_id": self.request_id, - "method": self.method, - "args": self.args - } + return {"request_id": self.request_id, "method": self.method, "args": self.args} def __repr__(self) -> str: """Provide a clean representation of the control request.""" @@ -449,42 +441,45 @@ def get_args(self) -> Dict[str, Any]: def is_control_request(d: dict) -> bool: """ Check if a dictionary represents a valid ControlRequest. - + Args: d: Dictionary to check - + Returns: bool: True if the dictionary contains the required fields for a ControlRequest """ - + # Check if all required fields are present and have correct types if not isinstance(d, dict): return False - + # Check field types if "request_id" not in d or not isinstance(d.get("request_id"), str): return False - + if "method" not in d or not isinstance(d.get("method"), str): return False - + # Args is optional, but if present should be a dict if "args" in d and not isinstance(d["args"], dict): return False - + return True + class ControlResponse: """ Response for control opeartions """ + def __init__( - self, - request_id: str, - error_code: int = 200, + self, + request_id: str, + error_code: int = 200, error_message: Optional[str] = None, result: Optional[dict] = None, - finished: bool = True) -> None: + finished: bool = True, + ) -> None: self.request_id = request_id self.finished = finished self.error_message = error_message @@ -498,7 +493,7 @@ def to_dict(self) -> dict: "finished": self.finished, "error_code": self.error_code, "error_message": self.error_message, - "result": self.result + "result": self.result, } @classmethod @@ -509,7 +504,7 @@ def from_dict(cls, d: dict): finished=d.get("finished", True), error_code=d.get("error_code", 200), error_message=d.get("error_message"), - result=d.get("result") + result=d.get("result"), ) def to_api_json_response(self) -> JSONResponse: @@ -519,7 +514,7 @@ def to_api_json_response(self) -> JSONResponse: "request_id": self.request_id, "status": status, "error_message": self.error_message, - "result": self.result + "result": self.result, } return JSONResponse(status_code=self.error_code, content=content) diff --git a/fastdeploy/engine/sched/resource_manager_v1.py b/fastdeploy/engine/sched/resource_manager_v1.py index 1cd570cb0b1..e595076e4df 100644 --- a/fastdeploy/engine/sched/resource_manager_v1.py +++ b/fastdeploy/engine/sched/resource_manager_v1.py @@ -258,7 +258,7 @@ def preempted_all(self): preempted_reqs = [] for i in range(len(self.running)): req = self.running.pop() - # txt2image: req.use_extend_tables is True, req can not be preempted. txt2image is not used in RL. + # txt2image: req.use_extend_tables is True, req can not be preempted. txt2image is not used in RL. if req.use_extend_tables: self.running.insert(0, req) continue @@ -281,8 +281,10 @@ def wait_worker_inflight_requests_finish(self, timeout=60): count += 1 time.sleep(0.001) if count >= timeout * 1000: - llm_logger.info(f"wait_inflight_requests_finish timeout after {timeout} seconds, " - f"still {len(self.to_be_rescheduled_request_id_set)} requests running") + llm_logger.info( + f"wait_inflight_requests_finish timeout after {timeout} seconds, " + f"still {len(self.to_be_rescheduled_request_id_set)} requests running" + ) def _trigger_preempt(self, request, num_new_blocks, preempted_reqs, scheduled_reqs): """ @@ -1269,12 +1271,14 @@ def update_metrics(self): main_process_metrics.num_requests_waiting.set(num_tasks - len(self.running)) def log_status(self): - llm_logger.info(f"ResourceManagerV1( " - f"waiting={len(self.waiting)}, " - f"running={len(self.running)}, " - f"preempted={len(self.to_be_rescheduled_request_id_set)}, " - f"tasks_list={self.tasks_list}, " - f"stop_flags={self.stop_flags}, " - f"req_dict={self.req_dict}, " - f"requests={self.requests}, " - f")") + llm_logger.info( + f"ResourceManagerV1( " + f"waiting={len(self.waiting)}, " + f"running={len(self.running)}, " + f"preempted={len(self.to_be_rescheduled_request_id_set)}, " + f"tasks_list={self.tasks_list}, " + f"stop_flags={self.stop_flags}, " + f"req_dict={self.req_dict}, " + f"requests={self.requests}, " + f")" + ) diff --git a/fastdeploy/entrypoints/engine_client.py b/fastdeploy/entrypoints/engine_client.py index 41b10697387..5230cc2ed8f 100644 --- a/fastdeploy/entrypoints/engine_client.py +++ b/fastdeploy/entrypoints/engine_client.py @@ -14,6 +14,7 @@ # limitations under the License. """ +import asyncio import inspect import os import time @@ -21,7 +22,6 @@ import uuid from copy import copy from http import HTTPStatus -import asyncio import numpy as np from filelock import FileLock @@ -29,6 +29,7 @@ import fastdeploy.metrics.trace as tracing from fastdeploy import envs from fastdeploy.config import FDConfig +from fastdeploy.engine.request import ControlRequest, ControlResponse from fastdeploy.entrypoints.openai.utils import DealerConnectionManager from fastdeploy.envs import FD_SUPPORT_MAX_CONNECTIONS from fastdeploy.eplb.utils import RedundantExpertWorkload @@ -52,7 +53,6 @@ api_server_logger, to_tensor, ) -from fastdeploy.engine.request import ControlRequest, ControlResponse class EngineClient: @@ -531,7 +531,6 @@ async def run_control_method(self, request: ControlRequest): api_server_logger.error(f"Error Run Control Method: {error_response}") return error_response - def is_workers_alive(self): """ Check the health of the model server by checking whether all workers are alive. diff --git a/fastdeploy/entrypoints/openai/api_server.py b/fastdeploy/entrypoints/openai/api_server.py index cb5faf51242..7b312444b4f 100644 --- a/fastdeploy/entrypoints/openai/api_server.py +++ b/fastdeploy/entrypoints/openai/api_server.py @@ -20,9 +20,9 @@ import threading import time import traceback +import uuid from collections.abc import AsyncGenerator from contextlib import asynccontextmanager -import uuid import uvicorn import zmq @@ -38,8 +38,8 @@ from fastdeploy.engine.args_utils import EngineArgs from fastdeploy.engine.async_llm import AsyncLLM from fastdeploy.engine.engine import LLMEngine -from fastdeploy.engine.request import ControlRequest from fastdeploy.engine.expert_service import ExpertService +from fastdeploy.engine.request import ControlRequest from fastdeploy.entrypoints.chat_utils import load_chat_template from fastdeploy.entrypoints.engine_client import EngineClient from fastdeploy.entrypoints.openai.middleware import AuthenticationMiddleware @@ -367,6 +367,7 @@ def ping(raw_request: Request) -> Response: """Ping check. Endpoint required for SageMaker""" return health(raw_request) + @app.post("/v1/pause") async def pause(request: Request) -> Response: # todo: support wait_for_inflight_requests(default False), clear_cache(default True) arguments @@ -375,6 +376,7 @@ async def pause(request: Request) -> Response: control_response = await app.state.engine_client.run_control_method(control_request) return control_response.to_api_json_response() + @app.post("/v1/resume") async def resume(request: Request) -> Response: request_id = f"control-{uuid.uuid4()}" @@ -382,6 +384,7 @@ async def resume(request: Request) -> Response: control_response = await app.state.engine_client.run_control_method(control_request) return control_response.to_api_json_response() + @app.get("/v1/is_paused") async def is_paused(request: Request) -> Response: request_id = f"control-{uuid.uuid4()}" @@ -389,27 +392,29 @@ async def is_paused(request: Request) -> Response: control_response = await app.state.engine_client.run_control_method(control_request) return control_response.to_api_json_response() + @app.post("/v1/update_weights") async def update_weights(request: Request) -> Response: request_id = f"control-{uuid.uuid4()}" # 从请求体中获取参数 request_data = await request.json() - + # 提取version和rsync_config参数 version = request_data.get("version") rsync_config = request_data.get("rsync_config") - + # 构建控制请求,传递参数 args = {} if version is not None: args["version"] = version if rsync_config is not None: args["rsync_config"] = rsync_config - + control_request = ControlRequest(request_id, "update_weights", args) control_response = await app.state.engine_client.run_control_method(control_request) return control_response.to_api_json_response() + def wrap_streaming_generator(original_generator: AsyncGenerator): """ Wrap an async generator to release the connection semaphore when the generator is finished. diff --git a/fastdeploy/rl/dynamic_weight_manager.py b/fastdeploy/rl/dynamic_weight_manager.py index c44fd6ec3d6..21eebb83b03 100644 --- a/fastdeploy/rl/dynamic_weight_manager.py +++ b/fastdeploy/rl/dynamic_weight_manager.py @@ -14,8 +14,8 @@ # limitations under the License. """ -import os import io +import os import time from multiprocessing.shared_memory import SharedMemory from typing import Any, Dict, List @@ -27,8 +27,10 @@ from fastdeploy.config import FDConfig from fastdeploy.inter_communicator import ModelWeightsStatus + def sync_weights_by_rdma(config, step, rank): from checkpoint_transfer.core import RDMAWeightsDownloader + downloader = RDMAWeightsDownloader(config) downloader.initialize() logger.info(f"Fetching weights for step:{step}, rank:{rank}...") @@ -98,9 +100,11 @@ def valid_parameters(old_state_dict, new_state_dict): logger.error(f"Invalid parameter: {key} not in new_state_dict") elif old_state_dict[key].shape != new_state_dict[key].shape: is_valid = False - logger.error(f"Invalid parameter: {key} shape mismatch, " - f"new shape:{new_state_dict[key].shape}, " - f"old shape:{old_state_dict[key].shape}") + logger.error( + f"Invalid parameter: {key} shape mismatch, " + f"new shape:{new_state_dict[key].shape}, " + f"old shape:{old_state_dict[key].shape}" + ) elif old_state_dict[key].dtype != new_state_dict[key].dtype: is_valid = False logger.error(f"Invalid parameter: {key} dtype mismatch") @@ -109,16 +113,20 @@ def valid_parameters(old_state_dict, new_state_dict): if rsync_config is None: rsync_config = self.fd_config.load_config.rsync_config if rsync_config is None or len(rsync_config) == "": - raise Exception(f"rsync config not set, please set it in 1) launch arguments '--rsync-config' " - f"or 2) interface arguments 'rsync_config'") + raise Exception( + "rsync config not set, please set it in 1) launch arguments '--rsync-config' " + "or 2) interface arguments 'rsync_config'" + ) if version is None or version == "": version = self.read_model_version_from_file() if version is None or version == "": - raise Exception(f"rsync model version not set, please set it in 1) {{model_version}}/version.txt " - f"or 2) interface arguments 'version'") + raise Exception( + "rsync model version not set, please set it in 1) {model_version}/version.txt " + "or 2) interface arguments 'version'" + ) - llm_logger.info(f"START update_weights_by_rdma, version:{version}, rsync_config:{rsync_config}") + logger.info(f"START update_weights_by_rdma, version:{version}, rsync_config:{rsync_config}") rank = self.local_rank sync_start = time.perf_counter() @@ -130,7 +138,7 @@ def valid_parameters(old_state_dict, new_state_dict): if not valid_parameters(old_state_dict, new_state_dict): logger.error("Invalid new_state_dict, update parameters failed") return - + update_start = time.perf_counter() for name, param in old_state_dict.items(): param.set_value(new_state_dict[name]) @@ -138,10 +146,17 @@ def valid_parameters(old_state_dict, new_state_dict): logger.info(f"params set value cost {update_cost:.2f} seconds") total_cost = time.perf_counter() - sync_start - logger.info(f"END update_weights_by_rdma, cost {total_cost:.2f} seconds", - f" version:{version}, rsync_config: {rsync_config}") - return {"sync_cost": sync_cost, "update_cost": update_cost, "total_cost": total_cost, "version": version, "rank": rank} - + logger.info( + f"END update_weights_by_rdma, cost {total_cost:.2f} seconds", + f" version:{version}, rsync_config: {rsync_config}", + ) + return { + "sync_cost": sync_cost, + "update_cost": update_cost, + "total_cost": total_cost, + "version": version, + "rank": rank, + } def update_parameters(self, pid: int = 0, restart_process_group=False) -> None: """Core method to update model parameters based on strategy.""" diff --git a/fastdeploy/worker/gpu_model_runner.py b/fastdeploy/worker/gpu_model_runner.py index 646fe260ab1..2a817029f85 100644 --- a/fastdeploy/worker/gpu_model_runner.py +++ b/fastdeploy/worker/gpu_model_runner.py @@ -20,7 +20,7 @@ import time from concurrent.futures import Future from threading import Thread -from typing import Any, List, Optional, cast, Dict +from typing import Any, Dict, List, Optional, cast import numpy as np import paddle @@ -1503,6 +1503,7 @@ def load_model(self) -> None: # 1.1 Load RL dynamic model if self.fd_config.load_config.dynamic_load_weight: from fastdeploy.rl.dynamic_weight_manager import DynamicWeightManager + self.dynamic_weight_manager = DynamicWeightManager(self.fd_config, self.model, self.local_rank) # 2. Load lora model diff --git a/fastdeploy/worker/gpu_worker.py b/fastdeploy/worker/gpu_worker.py index c98ad048593..ff974c1b369 100644 --- a/fastdeploy/worker/gpu_worker.py +++ b/fastdeploy/worker/gpu_worker.py @@ -16,7 +16,7 @@ import gc import time -from typing import List, Optional, Any, Dict +from typing import Any, Dict, List, Optional import paddle import pynvml @@ -224,4 +224,4 @@ def check_health(self) -> bool: def cal_theortical_kvcache(self) -> int: """Calculate the block memory required""" - return self.model_runner.cal_theortical_kvcache() \ No newline at end of file + return self.model_runner.cal_theortical_kvcache() diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index f31a7053488..0b32e2d0453 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -15,12 +15,12 @@ """ import argparse +import asyncio import json import os import time import traceback from typing import Tuple -import asyncio import numpy as np import paddle @@ -44,6 +44,7 @@ SpeculativeConfig, StructuredOutputsConfig, ) +from fastdeploy.engine.request import ControlRequest, ControlResponse from fastdeploy.eplb.async_expert_loader import ( MODEL_MAIN_NAME, REARRANGE_EXPERT_MAGIC_NUM, @@ -65,7 +66,6 @@ from fastdeploy.scheduler import SchedulerConfig from fastdeploy.utils import get_logger, optional_type from fastdeploy.worker.worker_base import WorkerBase -from fastdeploy.engine.request import ControlRequest, ControlResponse logger = get_logger("worker_process", "worker_process.log") @@ -660,7 +660,8 @@ def run_control_method(self, control_request: ControlRequest) -> None: def _send_control_response(self, control_response: ControlResponse): logger.info(f"Rank-{self.local_rank} put control output {control_response} to engine") - asyncio.run(self._ctrl_output.put(control_response, shm_threshold=100*1024*1024)) + asyncio.run(self._ctrl_output.put(control_response, shm_threshold=100 * 1024 * 1024)) + def parse_args(): """ From 9d2178fa86bbecbce436eda35cbf96c1df102409 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 15 Jan 2026 18:40:04 +0800 Subject: [PATCH 09/25] bug fix: update_weights support using default arguments --- fastdeploy/entrypoints/openai/api_server.py | 21 ++++++++------------- fastdeploy/rl/dynamic_weight_manager.py | 2 +- 2 files changed, 9 insertions(+), 14 deletions(-) diff --git a/fastdeploy/entrypoints/openai/api_server.py b/fastdeploy/entrypoints/openai/api_server.py index 7b312444b4f..32eb0cd378a 100644 --- a/fastdeploy/entrypoints/openai/api_server.py +++ b/fastdeploy/entrypoints/openai/api_server.py @@ -396,19 +396,14 @@ async def is_paused(request: Request) -> Response: @app.post("/v1/update_weights") async def update_weights(request: Request) -> Response: request_id = f"control-{uuid.uuid4()}" - # 从请求体中获取参数 - request_data = await request.json() - - # 提取version和rsync_config参数 - version = request_data.get("version") - rsync_config = request_data.get("rsync_config") - - # 构建控制请求,传递参数 - args = {} - if version is not None: - args["version"] = version - if rsync_config is not None: - args["rsync_config"] = rsync_config + + # 兼容无参数传入的情况 - 简洁写法 + request_data = await request.json() if await request.body() else {} + + # 提取并过滤有效参数 + args = { + key: value for key, value in request_data.items() if key in ("version", "rsync_config") and value is not None + } control_request = ControlRequest(request_id, "update_weights", args) control_response = await app.state.engine_client.run_control_method(control_request) diff --git a/fastdeploy/rl/dynamic_weight_manager.py b/fastdeploy/rl/dynamic_weight_manager.py index 21eebb83b03..149c18ea68b 100644 --- a/fastdeploy/rl/dynamic_weight_manager.py +++ b/fastdeploy/rl/dynamic_weight_manager.py @@ -147,7 +147,7 @@ def valid_parameters(old_state_dict, new_state_dict): total_cost = time.perf_counter() - sync_start logger.info( - f"END update_weights_by_rdma, cost {total_cost:.2f} seconds", + f"END update_weights_by_rdma, cost {total_cost:.2f} seconds" f" version:{version}, rsync_config: {rsync_config}", ) return { From 2b0b7fc73cb216e3497d3f3a3b7894281e742d22 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 15 Jan 2026 19:11:57 +0800 Subject: [PATCH 10/25] fix typo --- fastdeploy/config.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fastdeploy/config.py b/fastdeploy/config.py index 6a58874dd9a..dbd124026fa 100644 --- a/fastdeploy/config.py +++ b/fastdeploy/config.py @@ -1268,7 +1268,7 @@ def print(self): """ Print all configuration information. """ - i.info("EPLB Configuration Information :") + logger.info("EPLB Configuration Information :") for k, v in self.__dict__.items(): logger.info("{:<20}:{:<6}{}".format(k, "", v)) logger.info("=============================================================") From 85ce0742eb9a37aa90b4483399e42b74374886de Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 15 Jan 2026 19:36:33 +0800 Subject: [PATCH 11/25] typo fix --- fastdeploy/engine/request.py | 17 ++++++++--------- fastdeploy/worker/worker_process.py | 19 +++++++++---------- 2 files changed, 17 insertions(+), 19 deletions(-) diff --git a/fastdeploy/engine/request.py b/fastdeploy/engine/request.py index d567145c4cd..8e826e57f2b 100644 --- a/fastdeploy/engine/request.py +++ b/fastdeploy/engine/request.py @@ -17,6 +17,7 @@ from __future__ import annotations import time +import traceback from dataclasses import asdict, dataclass, fields from enum import Enum from typing import Any, Dict, Generic, Optional, Union @@ -204,14 +205,12 @@ def from_dict(cls, d: dict): # if mm_positions is not of type ImagePosition, convert to ImagePosition try: for i, mm_pos in enumerate(d["multimodal_inputs"]["mm_positions"]): - if not isinstance(mm_pos, ImagePosition): - if not isinstance(mm_pos, dict): - raise ValueError(f"Invalid mm_positions format at index {i}") - d["multimodal_inputs"]["mm_positions"][i] = ImagePosition(**mm_pos) - except (ValueError, TypeError, KeyError) as e: + d["multimodal_inputs"]["mm_positions"][i] = ( + ImagePosition(**mm_pos) if not isinstance(mm_pos, ImagePosition) else mm_pos + ) + except Exception as e: data_processor_logger.error( - f"Convert mm_positions to ImagePosition failed - {type(e).__name__}: {e}\n" - f"Input data: {d['multimodal_inputs']['mm_positions']}" + f"Convert mm_positions to ImagePosition error: {e}, {str(traceback.format_exc())}" ) raise return cls( @@ -965,8 +964,8 @@ class PoolingRequestOutput(Generic[_O]): prompt_token_ids: list[int] finished: bool metrics: Optional[RequestMetrics] = (None,) - error_code: Optional[int] = 200 - error_msg: Optional[str] = None + error_code: Optional[int] = (200,) + error_msg: Optional[str] = (None,) def __repr__(self): return ( diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 0b32e2d0453..8fc7a320a70 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -643,24 +643,23 @@ def run_control_method(self, control_request: ControlRequest) -> None: handler = getattr(self.worker, method, None) if handler is None or not callable(handler): - error_result = ControlResponse(request_id, 400, f"unknown control method {method}") - self._send_control_response(error_result) + error_msg = f"Rank-{self.local_rank}: Unknown control method {method}" + error_result = ControlResponse(request_id, 400, error_msg) + asyncio.run(self._ctrl_output.put(error_result)) return try: result = handler(**kwargs) succ_result = ControlResponse(request_id, 200, "Success", result) - self._send_control_response(succ_result) - logger.info(f"Success run control request: {control_request}, response: {succ_result}") + logger.info( + f"Rank-{self.local_rank} Success run control request: {control_request}, response: {succ_result}" + ) + asyncio.run(self._ctrl_output.put(succ_result, shm_threshold=100 * 1024 * 1024)) except Exception as e: - error_msg = f"Failed run control method {method}: {str(e)}" + error_msg = f"Rank-{self.local_rank} Failed run control method {method}: {str(e)}" logger.info(f"{error_msg}\n{traceback.format_exc()}") error_result = ControlResponse(request_id, 500, error_msg) - self._send_control_response(error_result) - - def _send_control_response(self, control_response: ControlResponse): - logger.info(f"Rank-{self.local_rank} put control output {control_response} to engine") - asyncio.run(self._ctrl_output.put(control_response, shm_threshold=100 * 1024 * 1024)) + asyncio.run(self._ctrl_output.put(error_result)) def parse_args(): From 7691b4521b4a8d82ce5979e267ebe2efe772fe2c Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 15 Jan 2026 19:37:46 +0800 Subject: [PATCH 12/25] typo fix --- fastdeploy/engine/request.py | 1 - 1 file changed, 1 deletion(-) diff --git a/fastdeploy/engine/request.py b/fastdeploy/engine/request.py index 8e826e57f2b..893c26de297 100644 --- a/fastdeploy/engine/request.py +++ b/fastdeploy/engine/request.py @@ -212,7 +212,6 @@ def from_dict(cls, d: dict): data_processor_logger.error( f"Convert mm_positions to ImagePosition error: {e}, {str(traceback.format_exc())}" ) - raise return cls( request_id=d["request_id"], prompt=d.get("prompt"), From 6ef8f9741ff4c989a5110c689d99408e4bc43617 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 15 Jan 2026 20:05:33 +0800 Subject: [PATCH 13/25] typo fix --- fastdeploy/engine/common_engine.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index 59813e1df7a..8066dfba51f 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -1226,7 +1226,7 @@ def _control_pause(self, control_request: ControlRequest) -> dict | None: break if count >= timeout * 1000: error_msg = f"wait engine_worker_queue tasks empty timeout after {timeout} seconds, worker may Hanged" - self.llm_logger.info(error_msg) + self.llm_logger.error(error_msg) raise Exception(error_msg) running_reqs = self.resource_manager.preempted_all() if len(running_reqs) > 0: From c7b7899299562cd472dc4780d1a6042e704bfb91 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Mon, 19 Jan 2026 18:13:28 +0800 Subject: [PATCH 14/25] add unitest for control request/response, localscheduler.get_inflight_requests, resource_manager_v1.preempted_all --- tests/engine/test_control_request_response.py | 329 ++++++++++++++++++ tests/engine/test_resource_manager_v1.py | 107 ++++++ tests/scheduler/test_local_scheduler.py | 14 + 3 files changed, 450 insertions(+) create mode 100644 tests/engine/test_control_request_response.py create mode 100644 tests/engine/test_resource_manager_v1.py diff --git a/tests/engine/test_control_request_response.py b/tests/engine/test_control_request_response.py new file mode 100644 index 00000000000..6d5b973bf81 --- /dev/null +++ b/tests/engine/test_control_request_response.py @@ -0,0 +1,329 @@ +# 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 unittest +from unittest.mock import patch + +from fastapi.responses import JSONResponse + +from fastdeploy.engine.request import ControlRequest, ControlResponse + + +class TestControlRequest(unittest.TestCase): + """Test cases for ControlRequest class.""" + + def test_initialization_basic(self): + """Test basic initialization of ControlRequest.""" + request_id = "test_request_123" + method = "get_metrics" + + request = ControlRequest(request_id=request_id, method=method) + + self.assertEqual(request.request_id, request_id) + self.assertEqual(request.method, method) + self.assertEqual(request.args, {}) + + def test_initialization_with_args(self): + """Test initialization with arguments.""" + request_id = "test_request_456" + method = "reset_scheduler" + args = {"force": True, "timeout": 30} + + request = ControlRequest(request_id=request_id, method=method, args=args) + + self.assertEqual(request.request_id, request_id) + self.assertEqual(request.method, method) + self.assertEqual(request.args, args) + + def test_from_dict_basic(self): + """Test creating ControlRequest from dictionary (basic case).""" + data = {"request_id": "test_from_dict", "method": "status_check"} + + request = ControlRequest.from_dict(data) + + self.assertEqual(request.request_id, data["request_id"]) + self.assertEqual(request.method, data["method"]) + self.assertEqual(request.args, {}) + + def test_from_dict_with_args(self): + """Test creating ControlRequest from dictionary with arguments.""" + data = { + "request_id": "test_from_dict_args", + "method": "configure", + "args": {"max_requests": 100, "queue_timeout": 60}, + } + + request = ControlRequest.from_dict(data) + + self.assertEqual(request.request_id, data["request_id"]) + self.assertEqual(request.method, data["method"]) + self.assertEqual(request.args, data["args"]) + + def test_to_dict_basic(self): + """Test converting ControlRequest to dictionary (basic case).""" + request = ControlRequest(request_id="test_to_dict", method="health_check") + + result = request.to_dict() + + expected = {"request_id": "test_to_dict", "method": "health_check", "args": {}} + self.assertEqual(result, expected) + + def test_to_dict_with_args(self): + """Test converting ControlRequest to dictionary with arguments.""" + args = {"setting1": "value1", "setting2": 42} + request = ControlRequest(request_id="test_to_dict_args", method="update_settings", args=args) + + result = request.to_dict() + + expected = {"request_id": "test_to_dict_args", "method": "update_settings", "args": args} + self.assertEqual(result, expected) + + def test_get_method(self): + """Test get_method method.""" + method = "custom_operation" + request = ControlRequest(request_id="test", method=method) + + self.assertEqual(request.get_method(), method) + + def test_get_args(self): + """Test get_args method.""" + args = {"param1": "value1", "param2": 123} + request = ControlRequest(request_id="test", method="test", args=args) + + result_args = request.get_args() + + self.assertEqual(result_args, args) + # Ensure it returns a copy, not the original dict + self.assertIsNot(result_args, args) + + def test_is_control_request_valid(self): + """Test is_control_request method with valid data.""" + valid_data = [ + {"request_id": "test1", "method": "method1"}, + {"request_id": "test2", "method": "method2", "args": {}}, + {"request_id": "test3", "method": "method3", "args": {"key": "value"}}, + ] + + for data in valid_data: + with self.subTest(data=data): + self.assertTrue(ControlRequest.is_control_request(data)) + + def test_is_control_request_invalid(self): + """Test is_control_request method with invalid data.""" + invalid_data = [ + # Missing required fields + {"method": "test"}, # missing request_id + {"request_id": "test"}, # missing method + # Wrong field types + {"request_id": 123, "method": "test"}, # request_id not string + {"request_id": "test", "method": 456}, # method not string + {"request_id": "test", "method": "test", "args": "not_a_dict"}, # args not dict + # Not a dict + "not_a_dict", + 123, + None, + ] + + for data in invalid_data: + with self.subTest(data=data): + self.assertFalse(ControlRequest.is_control_request(data)) + + def test_repr_simple(self): + """Test __repr__ method in simple mode.""" + with patch("fastdeploy.envs.FD_DEBUG", False): + request = ControlRequest(request_id="test_repr", method="test_method") + repr_str = repr(request) + + self.assertIn("ControlRequest", repr_str) + self.assertIn("test_repr", repr_str) + self.assertIn("test_method", repr_str) + self.assertNotIn("args", repr_str) # Args not shown in simple mode + + def test_repr_debug_mode(self): + """Test __repr__ method in debug mode.""" + with patch("fastdeploy.envs.FD_DEBUG", True): + args = {"debug_param": "debug_value"} + request = ControlRequest(request_id="test_repr", method="test_method", args=args) + repr_str = repr(request) + + self.assertIn("ControlRequest", repr_str) + self.assertIn("test_repr", repr_str) + self.assertIn("test_method", repr_str) + self.assertIn("debug_param", repr_str) # Args shown in debug mode + + +class TestControlResponse(unittest.TestCase): + """Test cases for ControlResponse class.""" + + def test_initialization_basic(self): + """Test basic initialization of ControlResponse.""" + request_id = "test_response_123" + + response = ControlResponse(request_id=request_id) + + self.assertEqual(response.request_id, request_id) + self.assertEqual(response.error_code, 200) + self.assertIsNone(response.error_message) + self.assertIsNone(response.result) + self.assertTrue(response.finished) + + def test_initialization_with_all_params(self): + """Test initialization with all parameters.""" + request_id = "test_response_456" + error_code = 404 + error_message = "Not found" + result = {"data": "some_result"} + finished = False + + response = ControlResponse( + request_id=request_id, error_code=error_code, error_message=error_message, result=result, finished=finished + ) + + self.assertEqual(response.request_id, request_id) + self.assertEqual(response.error_code, error_code) + self.assertEqual(response.error_message, error_message) + self.assertEqual(response.result, result) + self.assertEqual(response.finished, finished) + + def test_initialization_error_cases(self): + """Test initialization with various error codes.""" + test_cases = [ + (200, None, True), # Success case + (400, "Bad Request", False), # Client error + (500, "Internal Error", True), # Server error + ] + + for error_code, error_message, finished in test_cases: + with self.subTest(error_code=error_code): + response = ControlResponse( + request_id="test", error_code=error_code, error_message=error_message, finished=finished + ) + + self.assertEqual(response.error_code, error_code) + self.assertEqual(response.error_message, error_message) + self.assertEqual(response.finished, finished) + + def test_from_dict_basic(self): + """Test creating ControlResponse from dictionary (basic case).""" + data = {"request_id": "test_from_dict"} + + response = ControlResponse.from_dict(data) + + self.assertEqual(response.request_id, data["request_id"]) + self.assertEqual(response.error_code, 200) + self.assertIsNone(response.error_message) + self.assertIsNone(response.result) + self.assertTrue(response.finished) + + def test_from_dict_with_all_fields(self): + """Test creating ControlResponse from dictionary with all fields.""" + data = { + "request_id": "test_from_dict_full", + "error_code": 500, + "error_message": "Test error", + "result": {"key": "value"}, + "finished": False, + } + + response = ControlResponse.from_dict(data) + + self.assertEqual(response.request_id, data["request_id"]) + self.assertEqual(response.error_code, data["error_code"]) + self.assertEqual(response.error_message, data["error_message"]) + self.assertEqual(response.result, data["result"]) + self.assertEqual(response.finished, data["finished"]) + + def test_to_dict_basic(self): + """Test converting ControlResponse to dictionary (basic case).""" + response = ControlResponse(request_id="test_to_dict") + + result = response.to_dict() + + expected = { + "request_id": "test_to_dict", + "finished": True, + "error_code": 200, + "error_message": None, + "result": None, + } + self.assertEqual(result, expected) + + def test_to_dict_with_all_fields(self): + """Test converting ControlResponse to dictionary with all fields.""" + response = ControlResponse( + request_id="test_to_dict_full", + error_code=400, + error_message="Validation failed", + result={"valid": False, "reason": "missing_field"}, + finished=False, + ) + + result = response.to_dict() + + expected = { + "request_id": "test_to_dict_full", + "finished": False, + "error_code": 400, + "error_message": "Validation failed", + "result": {"valid": False, "reason": "missing_field"}, + } + self.assertEqual(result, expected) + + def test_to_api_json_response_success(self): + """Test converting to JSONResponse for successful response.""" + result_data = {"metrics": {"cpu_usage": 0.5, "memory_used": 1024}} + response = ControlResponse(request_id="test_json_success", result=result_data) + + json_response = response.to_api_json_response() + + self.assertIsInstance(json_response, JSONResponse) + self.assertEqual(json_response.status_code, 200) + + content = json_response.body.decode("utf-8") + self.assertIn("success", content) + self.assertIn("test_json_success", content) + self.assertIn("cpu_usage", content) + + def test_to_api_json_response_error(self): + """Test converting to JSONResponse for error response.""" + response = ControlResponse(request_id="test_json_error", error_code=503, error_message="Service unavailable") + + json_response = response.to_api_json_response() + + self.assertIsInstance(json_response, JSONResponse) + self.assertEqual(json_response.status_code, 503) + + content = json_response.body.decode("utf-8") + self.assertIn("error", content) + self.assertIn("test_json_error", content) + self.assertIn("Service unavailable", content) + + def test_repr_method(self): + """Test __repr__ method.""" + response = ControlResponse( + request_id="test_repr", error_code=200, error_message=None, result={"data": "test"}, finished=True + ) + + repr_str = repr(response) + + # Check that all important fields are represented + self.assertIn("ControlResponse", repr_str) + self.assertIn("test_repr", repr_str) + self.assertIn("200", repr_str) + self.assertIn("test", repr_str) # from result + self.assertIn("True", repr_str) # finished flag + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/engine/test_resource_manager_v1.py b/tests/engine/test_resource_manager_v1.py new file mode 100644 index 00000000000..0548a912cf2 --- /dev/null +++ b/tests/engine/test_resource_manager_v1.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 os +import unittest +from unittest.mock import Mock + +from fastdeploy.engine.args_utils import EngineArgs +from fastdeploy.engine.request import Request, RequestStatus +from fastdeploy.engine.sched.resource_manager_v1 import ResourceManagerV1 + +MODEL_NAME = os.getenv("MODEL_PATH", "/path/to/models") + "/ERNIE-4.5-0.3B-Paddle" + + +class TestResourceManagerV1(unittest.TestCase): + """Test cases for ResourceManagerV1.""" + + def setUp(self): + """Set up test fixtures.""" + engine_args = EngineArgs( + model=MODEL_NAME, + max_model_len=8192, + tensor_parallel_size=1, + engine_worker_queue_port=int(os.getenv("FD_ENGINE_QUEUE_PORT", "6778")), + cache_queue_port=int(os.getenv("FD_CACHE_QUEUE_PORT", "6779")), + ) + # Create and start the engine service + mock_config = engine_args.create_engine_config() + + self.manager = ResourceManagerV1( + max_num_seqs=4, + config=mock_config, + tensor_parallel_size=1, + splitwise_role="mixed", + local_data_parallel_id=0, + ) + + print("3") + # Mock cache manager + self.manager.cache_manager = Mock() + self.manager.cache_manager.free_blocks = Mock() + + def tearDown(self) -> None: + self.manager.need_block_num_signal.clear() + + def test_preempted_all_with_no_running_requests(self): + """Test preempted_all with no running requests.""" + print("hello") + self.assertEqual(len(self.manager.running), 0) + preempted_reqs = self.manager.preempted_all() + self.assertEqual(len(preempted_reqs), 0) + print("world") + + def test_preempted_all_with_normal_requests(self): + """Test preempted_all with normal running requests.""" + # Add mock running requests + req1 = Mock(spec=Request) + req1.request_id = "req1" + req1.use_extend_tables = False + req1.status = RequestStatus.RUNNING + req1.block_tables = [1, 2, 3] + req1.num_cached_blocks = 0 + req1.idx = 0 + + req2 = Mock(spec=Request) + req2.request_id = "req2" + req2.use_extend_tables = False + req2.status = RequestStatus.RUNNING + req2.block_tables = [4, 5] + req2.num_cached_blocks = 0 + req2.idx = 1 + + self.manager.running = [req1, req2] + + preempted_reqs = self.manager.preempted_all() + + # Verify + self.assertEqual(len(preempted_reqs), 2) + self.assertEqual(preempted_reqs[0].request_id, "req2") + self.assertEqual(preempted_reqs[1].request_id, "req1") + + # Verify request status changed + self.assertEqual(req1.status, RequestStatus.PREEMPTED) + self.assertEqual(req2.status, RequestStatus.PREEMPTED) + + # Verify added to to_be_rescheduled_request_id_set + self.assertIn("req1", self.manager.to_be_rescheduled_request_id_set) + self.assertIn("req2", self.manager.to_be_rescheduled_request_id_set) + + self.assertEqual(len(self.manager.running), 0) + self.assertEqual(len(self.manager.waiting), 0) + self.assertEqual(len(self.manager.to_be_rescheduled_request_id_set), 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/scheduler/test_local_scheduler.py b/tests/scheduler/test_local_scheduler.py index 616e20a13ea..48ef2844a09 100644 --- a/tests/scheduler/test_local_scheduler.py +++ b/tests/scheduler/test_local_scheduler.py @@ -246,6 +246,20 @@ def test_put_requests_duplicate_handling(self): # Verify only one request exists in scheduler self.assertEqual(len(self.scheduler.requests), 1) + def test_get_inflight_requests(self): + """Test getting inflight requests.""" + # Add some requests + requests = [self.mock_request_1, self.mock_request_2] + self.scheduler.put_requests(requests) + + # Get inflight requests + inflight_requests = self.scheduler.get_inflight_requests() + + # Verify correct requests are returned + self.assertEqual(len(inflight_requests), len(requests)) + for req in inflight_requests: + self.assertIn(req, requests) + def test_put_requests_max_size_limit(self): """Test that max size limit is enforced.""" # Create scheduler with small max size From d6a4f6703d60babdbc32921dd1ad71591e6d7d84 Mon Sep 17 00:00:00 2001 From: wangyifei Date: Mon, 19 Jan 2026 18:39:48 +0800 Subject: [PATCH 15/25] add "rsync" to LoadConfig.load_strategy Literal type hints Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- fastdeploy/config.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/fastdeploy/config.py b/fastdeploy/config.py index f03a377da4b..6fb967142fc 100644 --- a/fastdeploy/config.py +++ b/fastdeploy/config.py @@ -1177,7 +1177,7 @@ def __init__( ): self.load_choices: Union[str, LoadChoices] = LoadChoices.DEFAULT.value self.dynamic_load_weight: bool = False - self.load_strategy: Optional[Literal["ipc", "ipc_snapshot", "meta", "normal"]] = "normal" + self.load_strategy: Optional[Literal["ipc", "ipc_snapshot", "meta", "normal", "rsync"]] = "normal" self.rsync_config: Optional[Dict[str, Any]] = None for key, value in args.items(): if hasattr(self, key): From 07787c655c484ac7e155713fbaa24053307d0685 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Mon, 19 Jan 2026 18:59:34 +0800 Subject: [PATCH 16/25] typo fix --- fastdeploy/engine/request.py | 2 +- fastdeploy/entrypoints/engine_client.py | 1 - fastdeploy/rl/dynamic_weight_manager.py | 4 ++-- fastdeploy/worker/worker_process.py | 2 +- tests/engine/test_resource_manager_v1.py | 3 --- 5 files changed, 4 insertions(+), 8 deletions(-) diff --git a/fastdeploy/engine/request.py b/fastdeploy/engine/request.py index 36dce95c3c7..bf551ae636c 100644 --- a/fastdeploy/engine/request.py +++ b/fastdeploy/engine/request.py @@ -469,7 +469,7 @@ def is_control_request(d: dict) -> bool: class ControlResponse: """ - Response for control opeartions + Response for control operations """ def __init__( diff --git a/fastdeploy/entrypoints/engine_client.py b/fastdeploy/entrypoints/engine_client.py index 61ed9fa1130..1662d0b771d 100644 --- a/fastdeploy/entrypoints/engine_client.py +++ b/fastdeploy/entrypoints/engine_client.py @@ -526,7 +526,6 @@ async def run_control_method(self, request: ControlRequest): dealer.write([b"", request_id.encode("utf-8")]) try: response = await asyncio.wait_for(response_queue.get(), timeout=600) - print(response) response = ControlResponse.from_dict(response[0]) api_server_logger.info(f"End Run Control Method: {response}") return response diff --git a/fastdeploy/rl/dynamic_weight_manager.py b/fastdeploy/rl/dynamic_weight_manager.py index 0ce47d9b023..27032d2d3f2 100644 --- a/fastdeploy/rl/dynamic_weight_manager.py +++ b/fastdeploy/rl/dynamic_weight_manager.py @@ -91,7 +91,7 @@ def _capture_model_state(self): logger.info(f"Model param: {name}, shape={param.shape}, dtype={param.dtype}") self.state_dict[name] = param - def update_weights_by_rdma(self, version: str = None, rsync_config: dict[str, Any] = None): + def update_weights_by_rdma(self, version: str = None, rsync_config: Dict[str, Any] = None): def valid_parameters(old_state_dict, new_state_dict): is_valid = True for key in old_state_dict: @@ -112,7 +112,7 @@ def valid_parameters(old_state_dict, new_state_dict): if rsync_config is None: rsync_config = self.fd_config.load_config.rsync_config - if rsync_config is None or len(rsync_config) == "": + if rsync_config is None or len(rsync_config) == 0: raise Exception( "rsync config not set, please set it in 1) launch arguments '--rsync-config' " "or 2) interface arguments 'rsync_config'" diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index f5ecccc0cbf..78d8c474318 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -693,7 +693,7 @@ def run_control_method(self, control_request: ControlRequest) -> None: asyncio.run(self._ctrl_output.put(succ_result, shm_threshold=100 * 1024 * 1024)) except Exception as e: error_msg = f"Rank-{self.local_rank} Failed run control method {method}: {str(e)}" - logger.info(f"{error_msg}\n{traceback.format_exc()}") + logger.error(f"{error_msg}\n{traceback.format_exc()}") error_result = ControlResponse(request_id, 500, error_msg) asyncio.run(self._ctrl_output.put(error_result)) diff --git a/tests/engine/test_resource_manager_v1.py b/tests/engine/test_resource_manager_v1.py index 0548a912cf2..0031a2e4f69 100644 --- a/tests/engine/test_resource_manager_v1.py +++ b/tests/engine/test_resource_manager_v1.py @@ -46,7 +46,6 @@ def setUp(self): local_data_parallel_id=0, ) - print("3") # Mock cache manager self.manager.cache_manager = Mock() self.manager.cache_manager.free_blocks = Mock() @@ -56,11 +55,9 @@ def tearDown(self) -> None: def test_preempted_all_with_no_running_requests(self): """Test preempted_all with no running requests.""" - print("hello") self.assertEqual(len(self.manager.running), 0) preempted_reqs = self.manager.preempted_all() self.assertEqual(len(preempted_reqs), 0) - print("world") def test_preempted_all_with_normal_requests(self): """Test preempted_all with normal running requests.""" From 48b664c1cf960fba2096778bbe4c4f59d5380a14 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Mon, 19 Jan 2026 20:00:43 +0800 Subject: [PATCH 17/25] typo fix --- fastdeploy/engine/common_engine.py | 66 +++++++++++++++++------------- 1 file changed, 38 insertions(+), 28 deletions(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index 4d6187676ed..d95876b1430 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -1181,15 +1181,19 @@ def _insert_zmq_task_to_scheduler(self): def run_control_method(self, control_req: ControlRequest): """ - Execute control methods for engine management using dynamic method invocation. + Execute control method, process control request and return response. + + This method is responsible for handling control requests, calling the corresponding + handler function based on the method name in the request. If the method doesn't exist + or is not callable, it returns an error response; otherwise executes the method and + returns a success response. Args: - control_req: ControlRequest instance containing method name and arguments + control_req (ControlRequest): Control request object containing request ID, + method name and parameters. - Usage: - - Control request with method "get_metrics" will call self._control_get_metrics(args) - - Method names are automatically mapped to handler methods with prefix "_control_" - - If no handler exists, returns error with available methods + Returns: + None: No return value, sends ControlResponse through send_response_server. """ method = control_req.get_method() request_id = control_req.request_id @@ -1216,17 +1220,19 @@ def run_control_method(self, control_req: ControlRequest): error_result = ControlResponse(request_id, 500, error_msg) self.send_response_server.send_response(request_id, [error_result]) - def _control_pause(self, control_request: ControlRequest) -> dict | None: - """暂停请求生成 - + def _control_pause(self, control_request: ControlRequest): + """Pauses the LLM engine and aborts all running/inflight requests. Args: - args: 控制参数字典,暂停相关的配置参数 + control_request: The control request containing pause command + + Raises: + Exception: If pause is not supported in current configuration + Exception: If engine worker queue cleanup times out Returns: - tuple: (error_code, error_msg) 元组 - - error_code: 错误代码,0表示成功,非0表示失败 - - error_msg: 错误信息,成功时为空字符串 + None """ + if not envs.ENABLE_V1_KVCACHE_SCHEDULER: raise Exception("pause only supported in ENABLE_V1_KVCACHE_SCHEDULER") if self.cfg.scheduler_config.name != "local": @@ -1271,45 +1277,49 @@ def _control_pause(self, control_request: ControlRequest) -> dict | None: self.resource_manager.cache_manager.reset() return None - def _control_resume(self, control_request: ControlRequest) -> dict | None: - """恢复暂停的请求生成 + def _control_resume(self, control_request: ControlRequest) -> Optional[dict]: + """Control function for resuming request generation. - Args: - args: 控制参数字典,恢复生成相关的配置参数 + This method resumes the paused request generation process by setting the pause flag + and notifying all waiting threads. It logs the start and end of the resume operation. - Returns: - dict | None: 返回结果字典或None,包含恢复操作的状态信息 + Args: + control_request: Control request object containing resume operation information """ self.llm_logger.info("START Resume Request Generation") with self._pause_cond: if not self.is_paused: self.llm_logger.info("Resume Request Generation: not paused.") + return None self.is_paused = False self._pause_cond.notify_all() self.llm_logger.info("END Resume Request Generation") return None def _control_is_paused(self, control_request: ControlRequest) -> bool: - """检查是否暂停了请求生成 + """ + Check if the LLM engine is in paused state. Args: - args: 控制参数字典,检查是否暂停相关的配置参数 + control_request: Control request object. Returns: - bool: 是否暂停了请求生成 + dict: Dictionary containing pause status information, {'is_paused': bool} """ self.llm_logger.info(f"LLM Engine request generation is paused: {self.is_paused}") with self._pause_cond: return {"is_paused": self.is_paused} - def _control_update_weights(self, control_request: ControlRequest) -> dict | None: - """更新模型权重 - + def _control_update_weights(self, control_request: ControlRequest) -> Optional[dict]: + """Update model weights Args: - args: 控制参数字典,更新权重相关的配置参数 + control_request: Control request object containing parameters for weight updates Returns: - dict | None: 返回结果字典或None,包含更新权重的操作结果信息 + Optional[dict]: Returns the result dictionary if update succeeds, None otherwise + + Raises: + Exception: Raised when the engine is not in paused state """ self.llm_logger.info("Update Model Weights") with self._pause_cond: @@ -1337,7 +1347,7 @@ def _call_worker(self, control_request: ControlRequest, timeout: int): raise Exception(f"Call Worker error: {response.error_message}") self.llm_logger.info(f"Call Worker Succeed: {output_queue.name} {response.result}") responses.append(response.result) - return responses + return {"worker_responses": responses} def _send_error_response(self, request_id, error_msg, error_code: int = 500): self.llm_logger.error( From 32ab88f582398d258877a49611e1439463fe8485 Mon Sep 17 00:00:00 2001 From: wangyifei Date: Mon, 19 Jan 2026 20:06:53 +0800 Subject: [PATCH 18/25] Apply suggestion from @Copilot Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- fastdeploy/engine/common_engine.py | 43 ++++++++++++++++++++++++------ 1 file changed, 35 insertions(+), 8 deletions(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index d95876b1430..d483e2cf49c 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -1329,26 +1329,53 @@ def _control_update_weights(self, control_request: ControlRequest) -> Optional[d raise Exception(error_msg) return self._call_worker(control_request, 60) - def _call_worker(self, control_request: ControlRequest, timeout: int): - request_id = control_request.request_id - self.engine_worker_queue.put_tasks(([control_request], 1)) + async def _wait_all_control_responses(self, request_id: str, timeout: int): + """Wait for control responses from all workers with a global timeout. + + This method concurrently waits for responses from all control workers + and enforces an overall timeout to avoid leaking pending tasks. + """ + timeout_ms = timeout * 1000 + # Create one get() coroutine per worker output queue + tasks = [output_queue.get(timeout=timeout_ms) for output_queue in self._ctrl_worker_output_queues] + + try: + results = await asyncio.wait_for( + asyncio.gather(*tasks, return_exceptions=True), + timeout=timeout, + ) + except asyncio.TimeoutError: + # Keep the error message consistent with previous behavior + raise Exception("Worker Update Weights Timeouted after 600s") responses = [] - for output_queue in self._ctrl_worker_output_queues: - msg = asyncio.run(output_queue.get(timeout=timeout * 1000)) # todo: fix timeout when tp > 1 + for output_queue, msg in zip(self._ctrl_worker_output_queues, results): + if isinstance(msg, Exception): + self.llm_logger.error(f"Call Worker Failed: {output_queue.name} {repr(msg)}") + raise Exception(f"Call Worker error: {repr(msg)}") if msg is None: + # Preserve original semantics when no message is received raise Exception("Worker Update Weights Timeouted after 600s") response: ControlResponse = msg.payload if response.request_id != request_id: - self.llm_logger.info(f"ignore old control response from worker:{output_queue.name} {response}") + self.llm_logger.info( + f"ignore old control response from worker:{output_queue.name} {response}" + ) continue if response.error_code != 200: - self.llm_logger.info(f"Call Worker Failed: {output_queue.name} {response.error_message}") + self.llm_logger.info( + f"Call Worker Failed: {output_queue.name} {response.error_message}" + ) raise Exception(f"Call Worker error: {response.error_message}") self.llm_logger.info(f"Call Worker Succeed: {output_queue.name} {response.result}") responses.append(response.result) - return {"worker_responses": responses} + return responses + def _call_worker(self, control_request: ControlRequest, timeout: int): + request_id = control_request.request_id + self.engine_worker_queue.put_tasks(([control_request], 1)) + # Use a single asyncio.run() to concurrently wait for all worker responses. + return asyncio.run(self._wait_all_control_responses(request_id, timeout)) def _send_error_response(self, request_id, error_msg, error_code: int = 500): self.llm_logger.error( f"Send error response to client, request_id: {request_id}, error_msg: {error_msg}, error_code: {error_code}" From 277c1b5174ac49772d90544d8e00bd62fa79b4ec Mon Sep 17 00:00:00 2001 From: mitu626 Date: Mon, 19 Jan 2026 20:13:16 +0800 Subject: [PATCH 19/25] check version/rsync params --- fastdeploy/entrypoints/engine_client.py | 1 + fastdeploy/entrypoints/openai/api_server.py | 38 ++++++++++++++++++--- 2 files changed, 34 insertions(+), 5 deletions(-) diff --git a/fastdeploy/entrypoints/engine_client.py b/fastdeploy/entrypoints/engine_client.py index 1662d0b771d..4d55f005b66 100644 --- a/fastdeploy/entrypoints/engine_client.py +++ b/fastdeploy/entrypoints/engine_client.py @@ -525,6 +525,7 @@ async def run_control_method(self, request: ControlRequest): dealer, response_queue = await self.connection_manager.get_connection(request_id) dealer.write([b"", request_id.encode("utf-8")]) try: + # todo: support user specified timeout. default 600s is enough for most control cases response = await asyncio.wait_for(response_queue.get(), timeout=600) response = ControlResponse.from_dict(response[0]) api_server_logger.info(f"End Run Control Method: {response}") diff --git a/fastdeploy/entrypoints/openai/api_server.py b/fastdeploy/entrypoints/openai/api_server.py index 5361d80cc77..0bf3822bf54 100644 --- a/fastdeploy/entrypoints/openai/api_server.py +++ b/fastdeploy/entrypoints/openai/api_server.py @@ -401,13 +401,41 @@ async def is_paused(request: Request) -> Response: async def update_weights(request: Request) -> Response: request_id = f"control-{uuid.uuid4()}" - # 兼容无参数传入的情况 - 简洁写法 request_data = await request.json() if await request.body() else {} - # 提取并过滤有效参数 - args = { - key: value for key, value in request_data.items() if key in ("version", "rsync_config") and value is not None - } + args = {} + + # Validate and extract version parameter + if "version" in request_data and request_data["version"] is not None: + if not isinstance(request_data["version"], str): + return JSONResponse( + status_code=400, + content={ + "error": "Invalid parameter type", + "message": "version must be a string" + } + ) + args["version"] = request_data["version"] + + # Validate and extract rsync_config parameter + if "rsync_config" in request_data and request_data["rsync_config"] is not None: + if not isinstance(request_data["rsync_config"], dict): + return JSONResponse( + status_code=400, + content={ + "error": "Invalid parameter type", + "message": "rsync_config must be a dictionary" + } + ) + if "etcd_server" not in request_data["rsync_config"]: + return JSONResponse( + status_code=400, + content={ + "error": "Invalid parameter type", + "message": "rsync_config must contain etcd_server" + } + ) + args["rsync_config"] = request_data["rsync_config"] control_request = ControlRequest(request_id, "update_weights", args) control_response = await app.state.engine_client.run_control_method(control_request) From 0ab17a64c3b4340c283d2344026d468b2f041667 Mon Sep 17 00:00:00 2001 From: wangyifei Date: Mon, 19 Jan 2026 20:18:19 +0800 Subject: [PATCH 20/25] add error log when version.txt not exists Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- fastdeploy/rl/dynamic_weight_manager.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/fastdeploy/rl/dynamic_weight_manager.py b/fastdeploy/rl/dynamic_weight_manager.py index 27032d2d3f2..310f42cb49c 100644 --- a/fastdeploy/rl/dynamic_weight_manager.py +++ b/fastdeploy/rl/dynamic_weight_manager.py @@ -355,9 +355,14 @@ def _update_shared_status(self, pid: int, status: int) -> None: def read_model_version_from_file(self): model_dir = self.fd_config.model_config.model - with open(os.path.join(model_dir, "version.txt")) as f: - version = f.read().strip() - return version + version_file = os.path.join(model_dir, "version.txt") + try: + with open(version_file, "r", encoding="utf-8") as f: + version = f.read().strip() + return version + except (FileNotFoundError, OSError, IOError) as e: + logger.error(f"Failed to read model version file '{version_file}': {e}") + return None @staticmethod def check_model_weights_status(model_weights_status, kv_cache_status, model_runner, pid, block): From d4f76908ab7d991849c3f6e20852ad26f5041e8a Mon Sep 17 00:00:00 2001 From: wangyifei Date: Mon, 19 Jan 2026 20:20:21 +0800 Subject: [PATCH 21/25] raise specified ValueError when paramters check failed Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- fastdeploy/rl/dynamic_weight_manager.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/fastdeploy/rl/dynamic_weight_manager.py b/fastdeploy/rl/dynamic_weight_manager.py index 310f42cb49c..c932b76ed47 100644 --- a/fastdeploy/rl/dynamic_weight_manager.py +++ b/fastdeploy/rl/dynamic_weight_manager.py @@ -136,8 +136,9 @@ def valid_parameters(old_state_dict, new_state_dict): old_state_dict = self.state_dict if not valid_parameters(old_state_dict, new_state_dict): - logger.error("Invalid new_state_dict, update parameters failed") - return + error_msg = "Invalid new_state_dict, update parameters failed" + logger.error(error_msg) + raise ValueError(error_msg) update_start = time.perf_counter() for name, param in old_state_dict.items(): From 8845b822525cd515cada8ff4b6bea7c1a55354db Mon Sep 17 00:00:00 2001 From: mitu626 Date: Mon, 19 Jan 2026 20:23:00 +0800 Subject: [PATCH 22/25] tp barrier after run_control_method --- fastdeploy/entrypoints/openai/api_server.py | 24 ++++++--------------- fastdeploy/worker/worker_process.py | 2 ++ 2 files changed, 9 insertions(+), 17 deletions(-) diff --git a/fastdeploy/entrypoints/openai/api_server.py b/fastdeploy/entrypoints/openai/api_server.py index 0bf3822bf54..86d99c9513d 100644 --- a/fastdeploy/entrypoints/openai/api_server.py +++ b/fastdeploy/entrypoints/openai/api_server.py @@ -404,36 +404,26 @@ async def update_weights(request: Request) -> Response: request_data = await request.json() if await request.body() else {} args = {} - + # Validate and extract version parameter if "version" in request_data and request_data["version"] is not None: if not isinstance(request_data["version"], str): return JSONResponse( - status_code=400, - content={ - "error": "Invalid parameter type", - "message": "version must be a string" - } + status_code=400, content={"error": "Invalid parameter type", "message": "version must be a string"} ) args["version"] = request_data["version"] - + # Validate and extract rsync_config parameter if "rsync_config" in request_data and request_data["rsync_config"] is not None: if not isinstance(request_data["rsync_config"], dict): return JSONResponse( - status_code=400, - content={ - "error": "Invalid parameter type", - "message": "rsync_config must be a dictionary" - } + status_code=400, + content={"error": "Invalid parameter type", "message": "rsync_config must be a dictionary"}, ) if "etcd_server" not in request_data["rsync_config"]: return JSONResponse( - status_code=400, - content={ - "error": "Invalid parameter type", - "message": "rsync_config must contain etcd_server" - } + status_code=400, + content={"error": "Invalid parameter type", "message": "rsync_config must contain etcd_server"}, ) args["rsync_config"] = request_data["rsync_config"] diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index 78d8c474318..ecbec79efa9 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -529,10 +529,12 @@ def event_loop_normal(self) -> None: max_occupied_batch_index = int(bsz) req_dicts.extend(req_dict) + # todo: run control request async if len(control_reqs) > 0: logger.info(f"Rank: {self.local_rank} received {len(control_reqs)} control request.") for control_req in control_reqs: self.run_control_method(control_req) + self._tp_barrier_wait() if tp_size > 1 else None # Count prefill requests in current batch num_prefill_requests = sum(1 for req in req_dicts if req.task_type == RequestType.PREFILL) From 96a85d18f4313ecced720deb5b9105871fce4689 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 22 Jan 2026 14:35:15 +0800 Subject: [PATCH 23/25] encode 'engine_worker_queue_port' to unique name of worker2engine fmq queue --- fastdeploy/engine/common_engine.py | 12 +++++------- fastdeploy/worker/worker_process.py | 3 ++- 2 files changed, 7 insertions(+), 8 deletions(-) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index d483e2cf49c..3806ee115f8 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -104,7 +104,8 @@ def __init__(self, cfg, start_queue=True, use_async_llm=False): tp_size = cfg.parallel_config.tensor_parallel_size dp_index = cfg.parallel_config.local_data_parallel_id for rank in range(tp_size): - name = f"ctrl_w2e_rank{rank+tp_size*dp_index}" + engine_worker_queue_port = self.cfg.parallel_config.local_engine_worker_queue_port + name = f"ctrl_w2e_rank{rank+tp_size*dp_index}_{engine_worker_queue_port}" self.llm_logger.info(f"Init Worker Control Output Queue: {name}(consumer)") self._ctrl_worker_output_queues.append(FMQ().queue(name, "consumer")) @@ -1358,14 +1359,10 @@ async def _wait_all_control_responses(self, request_id: str, timeout: int): raise Exception("Worker Update Weights Timeouted after 600s") response: ControlResponse = msg.payload if response.request_id != request_id: - self.llm_logger.info( - f"ignore old control response from worker:{output_queue.name} {response}" - ) + self.llm_logger.info(f"ignore old control response from worker:{output_queue.name} {response}") continue if response.error_code != 200: - self.llm_logger.info( - f"Call Worker Failed: {output_queue.name} {response.error_message}" - ) + self.llm_logger.info(f"Call Worker Failed: {output_queue.name} {response.error_message}") raise Exception(f"Call Worker error: {response.error_message}") self.llm_logger.info(f"Call Worker Succeed: {output_queue.name} {response.result}") responses.append(response.result) @@ -1376,6 +1373,7 @@ def _call_worker(self, control_request: ControlRequest, timeout: int): self.engine_worker_queue.put_tasks(([control_request], 1)) # Use a single asyncio.run() to concurrently wait for all worker responses. return asyncio.run(self._wait_all_control_responses(request_id, timeout)) + def _send_error_response(self, request_id, error_msg, error_code: int = 500): self.llm_logger.error( f"Send error response to client, request_id: {request_id}, error_msg: {error_msg}, error_code: {error_code}" diff --git a/fastdeploy/worker/worker_process.py b/fastdeploy/worker/worker_process.py index ecbec79efa9..1ab30b62b10 100644 --- a/fastdeploy/worker/worker_process.py +++ b/fastdeploy/worker/worker_process.py @@ -168,7 +168,8 @@ def __init__(self, fd_config: FDConfig, ranks: int = 1, local_rank: int = 0) -> self.max_chips_per_node = 16 if current_platform.is_iluvatar() else 8 def init_control(self): - queue_name = f"ctrl_w2e_rank{self.local_rank}" + engine_worker_queue_port = self.parallel_config.local_engine_worker_queue_port + queue_name = f"ctrl_w2e_rank{self.local_rank}_{engine_worker_queue_port}" logger.info(f"Init Control Output Queue: {queue_name}(producer)") self._ctrl_output = FMQ().queue(queue_name, "producer") From 76f10fd4e2d4e31e4988e1fb79d9bf88675429a0 Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 22 Jan 2026 17:22:46 +0800 Subject: [PATCH 24/25] typo fix --- fastdeploy/engine/common_engine.py | 1 + fastdeploy/engine/sched/resource_manager_v1.py | 1 + 2 files changed, 2 insertions(+) diff --git a/fastdeploy/engine/common_engine.py b/fastdeploy/engine/common_engine.py index df78d7d0cf7..6ead226e3c0 100644 --- a/fastdeploy/engine/common_engine.py +++ b/fastdeploy/engine/common_engine.py @@ -44,6 +44,7 @@ ControlResponse, Request, RequestOutput, + RequestStatus, RequestType, ) from fastdeploy.engine.resource_manager import ResourceManager diff --git a/fastdeploy/engine/sched/resource_manager_v1.py b/fastdeploy/engine/sched/resource_manager_v1.py index b858fde3951..586b142777e 100644 --- a/fastdeploy/engine/sched/resource_manager_v1.py +++ b/fastdeploy/engine/sched/resource_manager_v1.py @@ -16,6 +16,7 @@ import copy import threading +import time import traceback from collections import deque from collections.abc import Iterable From 0aaede5538b97fc2afdaeb339a315168cfb1757d Mon Sep 17 00:00:00 2001 From: mitu626 Date: Thu, 22 Jan 2026 17:33:30 +0800 Subject: [PATCH 25/25] typo fix --- fastdeploy/entrypoints/engine_client.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/fastdeploy/entrypoints/engine_client.py b/fastdeploy/entrypoints/engine_client.py index 48676589065..d2f30ccf531 100644 --- a/fastdeploy/entrypoints/engine_client.py +++ b/fastdeploy/entrypoints/engine_client.py @@ -30,7 +30,12 @@ import fastdeploy.metrics.trace as tracing from fastdeploy import envs from fastdeploy.config import FDConfig -from fastdeploy.engine.request import ControlRequest, ControlResponse, Request, RequestStatus +from fastdeploy.engine.request import ( + ControlRequest, + ControlResponse, + Request, + RequestStatus, +) from fastdeploy.entrypoints.openai.utils import DealerConnectionManager from fastdeploy.envs import FD_SUPPORT_MAX_CONNECTIONS from fastdeploy.eplb.utils import RedundantExpertWorkload