-
Notifications
You must be signed in to change notification settings - Fork 756
[RL] add pause, update_weights, resume interface for async RL #6052
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
31 commits
Select commit
Hold shift + click to select a range
d927adf
support dynamic run_control_request through zmq from apiserver to com…
mitu626 b1e1f0e
support pause/resume/is_paused/update_weights in apiserver->common_en…
mitu626 c877720
change /is_puased from HTTP POST method to GET method
mitu626 c0bf665
add pause、resume、is_paused implementation
mitu626 004c57f
support engine <==> worker communication(request&response)
mitu626 3f1115c
support sync weights through RDMA from checkpoint_transfer
mitu626 ae306bc
support specified version, rsync_config in update_weights rpc call
mitu626 1b1a348
add pause, update_weights, resume interface for async RL
mitu626 9d2178f
bug fix: update_weights support using default arguments
mitu626 2b0b7fc
fix typo
mitu626 85ce074
typo fix
mitu626 7691b45
typo fix
mitu626 ae2acfe
Merge branch 'develop' into rlhf
mitu626 6ef8f97
typo fix
mitu626 6f61173
Merge remote-tracking branch 'refs/remotes/origin/rlhf' into rlhf
mitu626 05d0617
Merge branch 'develop' into rlhf
mitu626 c7b7899
add unitest for control request/response, localscheduler.get_inflight…
mitu626 c94d8de
Merge remote-tracking branch 'refs/remotes/origin/rlhf' into rlhf
mitu626 f643e95
Merge branch 'develop' into rlhf
Jiang-Jia-Jun d6a4f67
add "rsync" to LoadConfig.load_strategy Literal type hints
mitu626 07787c6
typo fix
mitu626 48b664c
typo fix
mitu626 32ab88f
Apply suggestion from @Copilot
mitu626 277c1b5
check version/rsync params
mitu626 0ab17a6
add error log when version.txt not exists
mitu626 d4f7690
raise specified ValueError when paramters check failed
mitu626 8845b82
tp barrier after run_control_method
mitu626 96a85d1
encode 'engine_worker_queue_port' to unique name of worker2engine fmq…
mitu626 4a4b053
Merge branch 'develop' into rlhf
mitu626 76f10fd
typo fix
mitu626 0aaede5
typo fix
mitu626 File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -16,6 +16,7 @@ | |
|
|
||
| from __future__ import annotations | ||
|
|
||
| import asyncio | ||
| import copy | ||
| import json | ||
| import multiprocessing | ||
|
|
@@ -38,7 +39,14 @@ | |
| from tqdm import tqdm | ||
|
|
||
| import fastdeploy.metrics.trace as tracing | ||
| from fastdeploy.engine.request import Request, RequestOutput, RequestStatus, RequestType | ||
| from fastdeploy.engine.request import ( | ||
| ControlRequest, | ||
| ControlResponse, | ||
| Request, | ||
| RequestOutput, | ||
| RequestStatus, | ||
| 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 | ||
|
|
@@ -50,6 +58,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 | ||
|
|
@@ -89,6 +98,18 @@ 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): | ||
| 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")) | ||
|
|
||
| self.scheduler = cfg.scheduler_config.scheduler() | ||
| self.enable_decode_cache_task = envs.FD_ENABLE_CACHE_TASK == "1" | ||
|
|
||
|
|
@@ -762,6 +783,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()), | ||
|
|
@@ -926,6 +949,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) | ||
|
|
@@ -1069,6 +1094,17 @@ def _insert_zmq_task_to_scheduler(self): | |
| self.recv_request_server = ZmqIpcServer(name=self.api_server_pid, mode=zmq.PULL) | ||
| continue | ||
|
|
||
| if ControlRequest.is_control_request(data): | ||
| 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: | ||
| self.llm_logger.error( | ||
| f"Failed to process control request {data.get('request_id')}: " | ||
| f"{e}, {traceback.format_exc()}" | ||
| ) | ||
| continue | ||
|
|
||
| request, insert_task = data, [] | ||
| results: List[Tuple[str, Optional[str]]] = list() | ||
| if data: | ||
|
|
@@ -1100,6 +1136,13 @@ 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) | ||
|
|
@@ -1139,6 +1182,200 @@ def _insert_zmq_task_to_scheduler(self): | |
| f"traceback={traceback.format_exc()}" | ||
| ) | ||
|
|
||
| def run_control_method(self, control_req: ControlRequest): | ||
| """ | ||
| 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): Control request object containing request ID, | ||
| method name and parameters. | ||
|
|
||
| Returns: | ||
| None: No return value, sends ControlResponse through send_response_server. | ||
| """ | ||
| 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): | ||
| 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 | ||
|
|
||
| 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()}") | ||
| 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): | ||
| """Pauses the LLM engine and aborts all running/inflight requests. | ||
| 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: | ||
| 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": | ||
| 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.") | ||
|
mitu626 marked this conversation as resolved.
|
||
| self.is_paused = True | ||
|
|
||
| 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 | ||
| 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.error(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)) | ||
|
mitu626 marked this conversation as resolved.
|
||
| self.resource_manager.wait_worker_inflight_requests_finish(timeout=60) | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 这里是需要等待推理自然结束还是中止推理呢
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 因为在上一步将当前所有running请求都调度成了抢占请求,因此worker会在下一个step将所有正在推理请求按照抢占逻辑打断,这里虽然是等待,但其实不是等待正常推理结束,而是在等待worker执行抢占操作。 |
||
| # self.engine_worker_queue.clear_data() | ||
| self.token_processor.clear_data() | ||
| self.resource_manager.log_status() | ||
|
|
||
| # abort inflight requests to user | ||
| 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.") | ||
| self.scheduler.reset() | ||
|
|
||
| self.resource_manager.cache_manager.reset() | ||
| return None | ||
|
|
||
| def _control_resume(self, control_request: ControlRequest) -> Optional[dict]: | ||
| """Control function for resuming request generation. | ||
|
|
||
| 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. | ||
|
|
||
| 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.") | ||
|
mitu626 marked this conversation as resolved.
|
||
| 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: | ||
| control_request: Control request object. | ||
|
|
||
| Returns: | ||
| dict: Dictionary containing pause status information, {'is_paused': bool} | ||
| """ | ||
|
mitu626 marked this conversation as resolved.
|
||
| 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) -> Optional[dict]: | ||
| """Update model weights | ||
| Args: | ||
| control_request: Control request object containing parameters for weight updates | ||
|
|
||
| Returns: | ||
| 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: | ||
| if self.is_paused is False: | ||
| 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) | ||
|
|
||
| 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, 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}") | ||
| 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 _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}" | ||
|
|
@@ -1712,6 +1949,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}" | ||
|
|
||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.