CoolFace
Apppublic

q-future/Co-Instruct

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
29likes
controller.py298 linesDownload Raw Back to serve
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")