OpenGVLab/InternVL
510
1"""2A controller manages distributed workers.3It sends worker addresses to clients.4"""5import argparse6import dataclasses7import json8import re9import threading10import time11from enum import Enum, auto12from typing import List13 14import numpy as np15import requests16import uvicorn17from fastapi import FastAPI, Request18from fastapi.responses import StreamingResponse19from utils import build_logger, server_error_msg20 21CONTROLLER_HEART_BEAT_EXPIRATION = 3022logger = build_logger('controller', 'controller.log')23 24 25class DispatchMethod(Enum):26 LOTTERY = auto()27 SHORTEST_QUEUE = auto()28 29 @classmethod30 def from_str(cls, name):31 if name == 'lottery':32 return cls.LOTTERY33 elif name == 'shortest_queue':34 return cls.SHORTEST_QUEUE35 else:36 raise ValueError(f'Invalid dispatch method')37 38 39@dataclasses.dataclass40class WorkerInfo:41 model_names: List[str]42 speed: int43 queue_length: int44 check_heart_beat: bool45 last_heart_beat: str46 47 48def heart_beat_controller(controller):49 while True:50 time.sleep(CONTROLLER_HEART_BEAT_EXPIRATION)51 controller.remove_stable_workers_by_expiration()52 53 54class Controller:55 def __init__(self, dispatch_method: str):56 # Dict[str -> WorkerInfo]57 self.worker_info = {}58 self.dispatch_method = DispatchMethod.from_str(dispatch_method)59 60 self.heart_beat_thread = threading.Thread(61 target=heart_beat_controller, args=(self,))62 self.heart_beat_thread.start()63 64 logger.info('Init controller')65 66 def register_worker(self, worker_name: str, check_heart_beat: bool,67 worker_status: dict):68 if worker_name not in self.worker_info:69 logger.info(f'Register a new worker: {worker_name}')70 else:71 logger.info(f'Register an existing worker: {worker_name}')72 73 if not worker_status:74 worker_status = self.get_worker_status(worker_name)75 if not worker_status:76 return False77 78 self.worker_info[worker_name] = WorkerInfo(79 worker_status['model_names'], worker_status['speed'], worker_status['queue_length'],80 check_heart_beat, time.time())81 82 logger.info(f'Register done: {worker_name}, {worker_status}')83 return True84 85 def get_worker_status(self, worker_name: str):86 try:87 r = requests.post(worker_name + '/worker_get_status', timeout=5)88 except requests.exceptions.RequestException as e:89 logger.error(f'Get status fails: {worker_name}, {e}')90 return None91 92 if r.status_code != 200:93 logger.error(f'Get status fails: {worker_name}, {r}')94 return None95 96 return r.json()97 98 def remove_worker(self, worker_name: str):99 del self.worker_info[worker_name]100 101 def refresh_all_workers(self):102 old_info = dict(self.worker_info)103 self.worker_info = {}104 105 for w_name, w_info in old_info.items():106 if not self.register_worker(w_name, w_info.check_heart_beat, None):107 logger.info(f'Remove stale worker: {w_name}')108 109 def list_models(self):110 model_names = set()111 112 for w_name, w_info in self.worker_info.items():113 model_names.update(w_info.model_names)114 115 def extract_key(s):116 if 'Pro' in s:117 return 999118 match = re.match(r'InternVL2-(\d+)B', s)119 if match:120 return int(match.group(1))121 return -1122 123 def custom_sort_key(s):124 key = extract_key(s)125 # Return a tuple where -1 will ensure that non-matching items come last126 return (0 if key != -1 else 1, -key if key != -1 else s)127 128 sorted_list = sorted(list(model_names), key=custom_sort_key)129 return sorted_list130 131 def get_worker_address(self, model_name: str):132 if self.dispatch_method == DispatchMethod.LOTTERY:133 worker_names = []134 worker_speeds = []135 for w_name, w_info in self.worker_info.items():136 if model_name in w_info.model_names:137 worker_names.append(w_name)138 worker_speeds.append(w_info.speed)139 worker_speeds = np.array(worker_speeds, dtype=np.float32)140 norm = np.sum(worker_speeds)141 if norm < 1e-4:142 return ''143 worker_speeds = worker_speeds / norm144 if True: # Directly return address145 pt = np.random.choice(np.arange(len(worker_names)),146 p=worker_speeds)147 worker_name = worker_names[pt]148 return worker_name149 150 elif self.dispatch_method == DispatchMethod.SHORTEST_QUEUE:151 worker_names = []152 worker_qlen = []153 for w_name, w_info in self.worker_info.items():154 if model_name in w_info.model_names:155 worker_names.append(w_name)156 worker_qlen.append(w_info.queue_length / w_info.speed)157 if len(worker_names) == 0:158 return ''159 min_index = np.argmin(worker_qlen)160 w_name = worker_names[min_index]161 self.worker_info[w_name].queue_length += 1162 logger.info(f'names: {worker_names}, queue_lens: {worker_qlen}, ret: {w_name}')163 return w_name164 else:165 raise ValueError(f'Invalid dispatch method: {self.dispatch_method}')166 167 def receive_heart_beat(self, worker_name: str, queue_length: int):168 if worker_name not in self.worker_info:169 logger.info(f'Receive unknown heart beat. {worker_name}')170 return False171 172 self.worker_info[worker_name].queue_length = queue_length173 self.worker_info[worker_name].last_heart_beat = time.time()174 logger.info(f'Receive heart beat. {worker_name}')175 return True176 177 def remove_stable_workers_by_expiration(self):178 expire = time.time() - CONTROLLER_HEART_BEAT_EXPIRATION179 to_delete = []180 for worker_name, w_info in self.worker_info.items():181 if w_info.check_heart_beat and w_info.last_heart_beat < expire:182 to_delete.append(worker_name)183 184 for worker_name in to_delete:185 self.remove_worker(worker_name)186 187 def worker_api_generate_stream(self, params):188 worker_addr = self.get_worker_address(params['model'])189 if not worker_addr:190 logger.info(f"no worker: {params['model']}")191 ret = {192 'text': server_error_msg,193 'error_code': 2,194 }195 yield json.dumps(ret).encode() + b'\0'196 197 try:198 response = requests.post(worker_addr + '/worker_generate_stream',199 json=params, stream=True, timeout=5)200 for chunk in response.iter_lines(decode_unicode=False, delimiter=b'\0'):201 if chunk:202 yield chunk + b'\0'203 except requests.exceptions.RequestException as e:204 logger.info(f'worker timeout: {worker_addr}')205 ret = {206 'text': server_error_msg,207 'error_code': 3,208 }209 yield json.dumps(ret).encode() + b'\0'210 211 # Let the controller act as a worker to achieve hierarchical212 # management. This can be used to connect isolated sub networks.213 def worker_api_get_status(self):214 model_names = set()215 speed = 0216 queue_length = 0217 218 for w_name in self.worker_info:219 worker_status = self.get_worker_status(w_name)220 if worker_status is not None:221 model_names.update(worker_status['model_names'])222 speed += worker_status['speed']223 queue_length += worker_status['queue_length']224 225 return {226 'model_names': list(model_names),227 'speed': speed,228 'queue_length': queue_length,229 }230 231 232app = FastAPI()233 234 235@app.post('/register_worker')236async def register_worker(request: Request):237 data = await request.json()238 controller.register_worker(239 data['worker_name'], data['check_heart_beat'],240 data.get('worker_status', None))241 242 243@app.post('/refresh_all_workers')244async def refresh_all_workers():245 models = controller.refresh_all_workers()246 247 248@app.post('/list_models')249async def list_models():250 models = controller.list_models()251 return {'models': models}252 253 254@app.post('/get_worker_address')255async def get_worker_address(request: Request):256 data = await request.json()257 addr = controller.get_worker_address(data['model'])258 return {'address': addr}259 260 261@app.post('/receive_heart_beat')262async def receive_heart_beat(request: Request):263 data = await request.json()264 exist = controller.receive_heart_beat(265 data['worker_name'], data['queue_length'])266 return {'exist': exist}267 268 269@app.post('/worker_generate_stream')270async def worker_api_generate_stream(request: Request):271 params = await request.json()272 generator = controller.worker_api_generate_stream(params)273 return StreamingResponse(generator)274 275 276@app.post('/worker_get_status')277async def worker_api_get_status(request: Request):278 return controller.worker_api_get_status()279 280 281if __name__ == '__main__':282 parser = argparse.ArgumentParser()283 parser.add_argument('--host', type=str, default='0.0.0.0')284 parser.add_argument('--port', type=int, default=10075)285 parser.add_argument('--dispatch-method', type=str, choices=[286 'lottery', 'shortest_queue'], default='shortest_queue')287 args = parser.parse_args()288 logger.info(f'args: {args}')289 290 controller = Controller(args.dispatch_method)291 uvicorn.run(app, host=args.host, port=args.port, log_level='info')292 