q-future/Co-Instruct
29
1"""2A controller manages distributed workers.3It sends worker addresses to clients.4"""5import argparse6import asyncio7import dataclasses8from enum import Enum, auto9import json10import logging11import time12from typing import List, Union13import threading14 15from fastapi import FastAPI, Request16from fastapi.responses import StreamingResponse17import numpy as np18import requests19import uvicorn20 21from mplug_owl2.constants import CONTROLLER_HEART_BEAT_EXPIRATION22from mplug_owl2.utils import build_logger, server_error_msg23 24 25logger = build_logger("controller", "controller.log")26 27 28class DispatchMethod(Enum):29 LOTTERY = auto()30 SHORTEST_QUEUE = auto()31 32 @classmethod33 def from_str(cls, name):34 if name == "lottery":35 return cls.LOTTERY36 elif name == "shortest_queue":37 return cls.SHORTEST_QUEUE38 else:39 raise ValueError(f"Invalid dispatch method")40 41 42@dataclasses.dataclass43class WorkerInfo:44 model_names: List[str]45 speed: int46 queue_length: int47 check_heart_beat: bool48 last_heart_beat: str49 50 51def heart_beat_controller(controller):52 while True:53 time.sleep(CONTROLLER_HEART_BEAT_EXPIRATION)54 controller.remove_stable_workers_by_expiration()55 56 57class Controller:58 def __init__(self, dispatch_method: str):59 # Dict[str -> WorkerInfo]60 self.worker_info = {}61 self.dispatch_method = DispatchMethod.from_str(dispatch_method)62 63 self.heart_beat_thread = threading.Thread(64 target=heart_beat_controller, args=(self,))65 self.heart_beat_thread.start()66 67 logger.info("Init controller")68 69 def register_worker(self, worker_name: str, check_heart_beat: bool,70 worker_status: dict):71 if worker_name not in self.worker_info:72 logger.info(f"Register a new worker: {worker_name}")73 else:74 logger.info(f"Register an existing worker: {worker_name}")75 76 if not worker_status:77 worker_status = self.get_worker_status(worker_name)78 if not worker_status:79 return False80 81 self.worker_info[worker_name] = WorkerInfo(82 worker_status["model_names"], worker_status["speed"], worker_status["queue_length"],83 check_heart_beat, time.time())84 85 logger.info(f"Register done: {worker_name}, {worker_status}")86 return True87 88 def get_worker_status(self, worker_name: str):89 try:90 r = requests.post(worker_name + "/worker_get_status", timeout=5)91 except requests.exceptions.RequestException as e:92 logger.error(f"Get status fails: {worker_name}, {e}")93 return None94 95 if r.status_code != 200:96 logger.error(f"Get status fails: {worker_name}, {r}")97 return None98 99 return r.json()100 101 def remove_worker(self, worker_name: str):102 del self.worker_info[worker_name]103 104 def refresh_all_workers(self):105 old_info = dict(self.worker_info)106 self.worker_info = {}107 108 for w_name, w_info in old_info.items():109 if not self.register_worker(w_name, w_info.check_heart_beat, None):110 logger.info(f"Remove stale worker: {w_name}")111 112 def list_models(self):113 model_names = set()114 115 for w_name, w_info in self.worker_info.items():116 model_names.update(w_info.model_names)117 118 return list(model_names)119 120 def get_worker_address(self, model_name: str):121 if self.dispatch_method == DispatchMethod.LOTTERY:122 worker_names = []123 worker_speeds = []124 for w_name, w_info in self.worker_info.items():125 if model_name in w_info.model_names:126 worker_names.append(w_name)127 worker_speeds.append(w_info.speed)128 worker_speeds = np.array(worker_speeds, dtype=np.float32)129 norm = np.sum(worker_speeds)130 if norm < 1e-4:131 return ""132 worker_speeds = worker_speeds / norm133 if True: # Directly return address134 pt = np.random.choice(np.arange(len(worker_names)),135 p=worker_speeds)136 worker_name = worker_names[pt]137 return worker_name138 139 # Check status before returning140 while True:141 pt = np.random.choice(np.arange(len(worker_names)),142 p=worker_speeds)143 worker_name = worker_names[pt]144 145 if self.get_worker_status(worker_name):146 break147 else:148 self.remove_worker(worker_name)149 worker_speeds[pt] = 0150 norm = np.sum(worker_speeds)151 if norm < 1e-4:152 return ""153 worker_speeds = worker_speeds / norm154 continue155 return worker_name156 elif self.dispatch_method == DispatchMethod.SHORTEST_QUEUE:157 worker_names = []158 worker_qlen = []159 for w_name, w_info in self.worker_info.items():160 if model_name in w_info.model_names:161 worker_names.append(w_name)162 worker_qlen.append(w_info.queue_length / w_info.speed)163 if len(worker_names) == 0:164 return ""165 min_index = np.argmin(worker_qlen)166 w_name = worker_names[min_index]167 self.worker_info[w_name].queue_length += 1168 logger.info(f"names: {worker_names}, queue_lens: {worker_qlen}, ret: {w_name}")169 return w_name170 else:171 raise ValueError(f"Invalid dispatch method: {self.dispatch_method}")172 173 def receive_heart_beat(self, worker_name: str, queue_length: int):174 if worker_name not in self.worker_info:175 logger.info(f"Receive unknown heart beat. {worker_name}")176 return False177 178 self.worker_info[worker_name].queue_length = queue_length179 self.worker_info[worker_name].last_heart_beat = time.time()180 logger.info(f"Receive heart beat. {worker_name}")181 return True182 183 def remove_stable_workers_by_expiration(self):184 expire = time.time() - CONTROLLER_HEART_BEAT_EXPIRATION185 to_delete = []186 for worker_name, w_info in self.worker_info.items():187 if w_info.check_heart_beat and w_info.last_heart_beat < expire:188 to_delete.append(worker_name)189 190 for worker_name in to_delete:191 self.remove_worker(worker_name)192 193 def worker_api_generate_stream(self, params):194 worker_addr = self.get_worker_address(params["model"])195 if not worker_addr:196 logger.info(f"no worker: {params['model']}")197 ret = {198 "text": server_error_msg,199 "error_code": 2,200 }201 yield json.dumps(ret).encode() + b"\0"202 203 try:204 response = requests.post(worker_addr + "/worker_generate_stream",205 json=params, stream=True, timeout=5)206 for chunk in response.iter_lines(decode_unicode=False, delimiter=b"\0"):207 if chunk:208 yield chunk + b"\0"209 except requests.exceptions.RequestException as e:210 logger.info(f"worker timeout: {worker_addr}")211 ret = {212 "text": server_error_msg,213 "error_code": 3,214 }215 yield json.dumps(ret).encode() + b"\0"216 217 218 # Let the controller act as a worker to achieve hierarchical219 # management. This can be used to connect isolated sub networks.220 def worker_api_get_status(self):221 model_names = set()222 speed = 0223 queue_length = 0224 225 for w_name in self.worker_info:226 worker_status = self.get_worker_status(w_name)227 if worker_status is not None:228 model_names.update(worker_status["model_names"])229 speed += worker_status["speed"]230 queue_length += worker_status["queue_length"]231 232 return {233 "model_names": list(model_names),234 "speed": speed,235 "queue_length": queue_length,236 }237 238 239app = FastAPI()240 241 242@app.post("/register_worker")243async def register_worker(request: Request):244 data = await request.json()245 controller.register_worker(246 data["worker_name"], data["check_heart_beat"],247 data.get("worker_status", None))248 249 250@app.post("/refresh_all_workers")251async def refresh_all_workers():252 models = controller.refresh_all_workers()253 254 255@app.post("/list_models")256async def list_models():257 models = controller.list_models()258 return {"models": models}259 260 261@app.post("/get_worker_address")262async def get_worker_address(request: Request):263 data = await request.json()264 addr = controller.get_worker_address(data["model"])265 return {"address": addr}266 267 268@app.post("/receive_heart_beat")269async def receive_heart_beat(request: Request):270 data = await request.json()271 exist = controller.receive_heart_beat(272 data["worker_name"], data["queue_length"])273 return {"exist": exist}274 275 276@app.post("/worker_generate_stream")277async def worker_api_generate_stream(request: Request):278 params = await request.json()279 generator = controller.worker_api_generate_stream(params)280 return StreamingResponse(generator)281 282 283@app.post("/worker_get_status")284async def worker_api_get_status(request: Request):285 return controller.worker_api_get_status()286 287 288if __name__ == "__main__":289 parser = argparse.ArgumentParser()290 parser.add_argument("--host", type=str, default="localhost")291 parser.add_argument("--port", type=int, default=21001)292 parser.add_argument("--dispatch-method", type=str, choices=[293 "lottery", "shortest_queue"], default="shortest_queue")294 args = parser.parse_args()295 logger.info(f"args: {args}")296 297 controller = Controller(args.dispatch_method)298 uvicorn.run(app, host=args.host, port=args.port, log_level="info")