CoolFace
Apppublic

OpenGVLab/InternVL

sourceHugging Facemitupdated 2y agoView on Hugging Face
510likes
controller.py292 linesDownload Raw Back to root
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