CoolFace
Apppublic

OpenGVLab/InternVL

sourceHugging Facemitupdated 2y agoView on Hugging Face
510likes
model_worker.py553 linesDownload Raw Back to root
1# --------------------------------------------------------2# InternVL3# Copyright (c) 2024 OpenGVLab4# Licensed under The MIT License [see LICENSE for details]5# --------------------------------------------------------6 7"""8A model worker executes the model.9"""10import spaces11import os12import argparse13import asyncio14 15import json16import math17import threading18import time19import uuid20import traceback21from functools import partial22from threading import Thread23 24import requests25import torch26import torchvision.transforms as T27import uvicorn28from constants import IMAGENET_MEAN, IMAGENET_STD, WORKER_HEART_BEAT_INTERVAL29from fastapi import BackgroundTasks, FastAPI, Request30from fastapi.responses import StreamingResponse31from PIL import Image32from torchvision.transforms.functional import InterpolationMode33from transformers import AutoModel, AutoTokenizer, TextIteratorStreamer34from utils import (35    build_logger,36    pretty_print_semaphore,37    server_error_msg,38    load_image_from_base64,39)40 41 42worker_id = str(uuid.uuid4())[:6]43logger = build_logger("model_worker", f"model_worker_{worker_id}.log")44global_counter = 045model_semaphore = None46 47 48def build_transform(input_size):49    MEAN, STD = IMAGENET_MEAN, IMAGENET_STD50    transform = T.Compose(51        [52            T.Lambda(lambda img: img.convert("RGB") if img.mode != "RGB" else img),53            T.Resize((input_size, input_size), interpolation=InterpolationMode.BICUBIC),54            T.ToTensor(),55            T.Normalize(mean=MEAN, std=STD),56        ]57    )58    return transform59 60 61def find_closest_aspect_ratio(aspect_ratio, target_ratios, width, height, image_size):62    best_ratio_diff = float("inf")63    best_ratio = (1, 1)64    area = width * height65    for ratio in target_ratios:66        target_aspect_ratio = ratio[0] / ratio[1]67        ratio_diff = abs(aspect_ratio - target_aspect_ratio)68        if ratio_diff < best_ratio_diff:69            best_ratio_diff = ratio_diff70            best_ratio = ratio71        elif ratio_diff == best_ratio_diff:72            if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:73                best_ratio = ratio74    return best_ratio75 76 77def dynamic_preprocess(78    image, min_num=1, max_num=6, image_size=448, use_thumbnail=False79):80    orig_width, orig_height = image.size81    aspect_ratio = orig_width / orig_height82 83    # calculate the existing image aspect ratio84    target_ratios = set(85        (i, j)86        for n in range(min_num, max_num + 1)87        for i in range(1, n + 1)88        for j in range(1, n + 1)89        if i * j <= max_num and i * j >= min_num90    )91    target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])92 93    # find the closest aspect ratio to the target94    target_aspect_ratio = find_closest_aspect_ratio(95        aspect_ratio, target_ratios, orig_width, orig_height, image_size96    )97 98    # calculate the target width and height99    target_width = image_size * target_aspect_ratio[0]100    target_height = image_size * target_aspect_ratio[1]101    blocks = target_aspect_ratio[0] * target_aspect_ratio[1]102 103    # resize the image104    resized_img = image.resize((target_width, target_height))105    processed_images = []106    for i in range(blocks):107        box = (108            (i % (target_width // image_size)) * image_size,109            (i // (target_width // image_size)) * image_size,110            ((i % (target_width // image_size)) + 1) * image_size,111            ((i // (target_width // image_size)) + 1) * image_size,112        )113        # split the image114        split_img = resized_img.crop(box)115        processed_images.append(split_img)116    assert len(processed_images) == blocks117    if use_thumbnail and len(processed_images) != 1:118        thumbnail_img = image.resize((image_size, image_size))119        processed_images.append(thumbnail_img)120    return processed_images121 122 123def heart_beat_worker(controller):124    while True:125        time.sleep(WORKER_HEART_BEAT_INTERVAL)126        controller.send_heart_beat()127 128 129def split_model(model_name):130    device_map = {}131    world_size = torch.cuda.device_count()132    num_layers = {133        "InternVL2-8B": 32,134        "InternVL2-26B": 48,135        "InternVL2-40B": 60,136        "InternVL2-Llama3-76B": 80,137        "InternVL2-78B": 80,138        "InternVL2-Pro": 80,139    }[model_name]140    # Since the first GPU will be used for ViT, treat it as half a GPU.141    num_layers_per_gpu = math.ceil(num_layers / (world_size - 0.5))142    num_layers_per_gpu = [num_layers_per_gpu] * world_size143    num_layers_per_gpu[0] = math.ceil(num_layers_per_gpu[0] * 0.5)144    layer_cnt = 0145    for i, num_layer in enumerate(num_layers_per_gpu):146        for j in range(num_layer):147            device_map[f"language_model.model.layers.{layer_cnt}"] = i148            layer_cnt += 1149    device_map["vision_model"] = 0150    device_map["mlp1"] = 0151    device_map["language_model.model.tok_embeddings"] = 0152    device_map["language_model.model.embed_tokens"] = 0153    device_map["language_model.output"] = 0154    device_map["language_model.model.norm"] = 0155    device_map["language_model.lm_head"] = 0156    device_map[f"language_model.model.layers.{num_layers - 1}"] = 0157 158    return device_map159 160 161def multi_thread_infer(162    model, tokenizer, pixel_values, question, history, generation_config163):164    with torch.no_grad():165        thread = Thread(166            target=model.chat,167            kwargs=dict(168                tokenizer=tokenizer,169                pixel_values=pixel_values,170                question=question,171                history=history,172                return_history=False,173                generation_config=generation_config,174            ),175        )176        thread.start()177 178 179class ModelWorker:180    def __init__(181        self,182        controller_addr,183        worker_addr,184        worker_id,185        model_path,186        model_name,187        load_8bit,188        device,189        context_len=8192,190    ):191        self.controller_addr = controller_addr192        self.worker_addr = worker_addr193        self.worker_id = worker_id194        if model_path.endswith("/"):195            model_path = model_path[:-1]196        if model_name is None:197            model_paths = model_path.split("/")198            if model_paths[-1].startswith("checkpoint-"):199                self.model_name = model_paths[-2] + "_" + model_paths[-1]200            else:201                self.model_name = model_paths[-1]202        else:203            self.model_name = model_name204 205        self.import_flash_attn()206        logger.info(f"Loading the model {self.model_name} on worker {worker_id} ...")207        tokenizer = AutoTokenizer.from_pretrained(208            model_path, trust_remote_code=True, use_fast=False209        )210        tokens_to_keep = ["<box>", "</box>", "<ref>", "</ref>"]211        tokenizer.additional_special_tokens = [212            item213            for item in tokenizer.additional_special_tokens214            if item not in tokens_to_keep215        ]216        self.tokenizer = tokenizer217 218        if device == "auto":219            device_map = split_model(self.model_name)220            self.model = AutoModel.from_pretrained(221                model_path,222                load_in_8bit=load_8bit,223                torch_dtype=torch.bfloat16,224                device_map=device_map,225                trust_remote_code=True,226            ).eval()227        else:228            self.model = AutoModel.from_pretrained(229                model_path,230                load_in_8bit=load_8bit,231                torch_dtype=torch.bfloat16,232                trust_remote_code=True,233            ).eval()234        if not load_8bit and not device == "auto":235            self.model = self.model.cuda()236        self.load_8bit = load_8bit237        self.device = device238        self.model_path = model_path239        self.image_size = self.model.config.force_image_size240        self.context_len = context_len241        self.register_to_controller()242        self.heart_beat_thread = threading.Thread(243            target=heart_beat_worker, args=(self,)244        )245        self.heart_beat_thread.start()246 247    @spaces.GPU(duration=120)248    def import_flash_attn(self):249        try:250            import flash_attn251        except ImportError:252 253            def install_flash_attn():254                os.system(255                    "FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn==2.5.9.post1 --no-build-isolation"256                )257 258            install_flash_attn()259            # import flash_attn260 261    def reload_model(self):262        del self.model263        torch.cuda.empty_cache()264        if self.device == "auto":265            device_map = split_model(self.model_name)266            self.model = AutoModel.from_pretrained(267                self.model_path,268                load_in_8bit=self.load_8bit,269                torch_dtype=torch.bfloat16,270                device_map=device_map,271                trust_remote_code=True,272            ).eval()273        else:274            self.model = AutoModel.from_pretrained(275                self.model_path,276                load_in_8bit=self.load_8bit,277                torch_dtype=torch.bfloat16,278                trust_remote_code=True,279            ).eval()280        if not self.load_8bit and not self.device == "auto":281            self.model = self.model.cuda()282 283    def register_to_controller(self):284        logger.info("Register to controller")285 286        url = self.controller_addr + "/register_worker"287        data = {288            "worker_name": self.worker_addr,289            "check_heart_beat": True,290            "worker_status": self.get_status(),291        }292        r = requests.post(url, json=data)293        assert r.status_code == 200294 295    def send_heart_beat(self):296        logger.info(297            f"Send heart beat. Models: {[self.model_name]}. "298            f"Semaphore: {pretty_print_semaphore(model_semaphore)}. "299            f"global_counter: {global_counter}"300        )301 302        url = self.controller_addr + "/receive_heart_beat"303 304        while True:305            try:306                ret = requests.post(307                    url,308                    json={309                        "worker_name": self.worker_addr,310                        "queue_length": self.get_queue_length(),311                    },312                    timeout=5,313                )314                exist = ret.json()["exist"]315                break316            except requests.exceptions.RequestException as e:317                logger.error(f"heart beat error: {e}")318            time.sleep(5)319 320        if not exist:321            self.register_to_controller()322 323    def get_queue_length(self):324        if model_semaphore is None:325            return 0326        else:327            return (328                args.limit_model_concurrency329                - model_semaphore._value330                + (331                    len(model_semaphore._waiters)332                    if model_semaphore._waiters is not None333                    else 0334                )335            )336 337    def get_status(self):338        return {339            "model_names": [self.model_name],340            "speed": 1,341            "queue_length": self.get_queue_length(),342        }343 344    def generate_stream(self, params):345        system_message = params["prompt"][0]["content"]346        send_messages = params["prompt"][1:]347        max_input_tiles = params["max_input_tiles"]348        temperature = params["temperature"]349        top_p = params["top_p"]350        max_new_tokens = params["max_new_tokens"]351        repetition_penalty = params["repetition_penalty"]352        do_sample = True if temperature > 0.0 else False353 354        global_image_cnt = 0355        history, pil_images, max_input_tile_list = [], [], []356        for message in send_messages:357            if message["role"] == "user":358                prefix = ""359                if "image" in message:360                    max_input_tile_temp = []361                    for image_str in message["image"]:362                        pil_images.append(load_image_from_base64(image_str))363                        prefix += f"Image-{global_image_cnt + 1}: <image>\n\n"364                        global_image_cnt += 1365                        max_input_tile_temp.append(366                            max(1, max_input_tiles // len(message["image"]))367                        )368                    if len(max_input_tile_temp) > 0:369                        max_input_tile_list.append(max_input_tile_temp)370                content = prefix + message["content"]371                history.append(372                    [373                        content,374                    ]375                )376            else:377                history[-1].append(message["content"])378        question, history = history[-1][0], history[:-1]379 380        if global_image_cnt == 1:381            question = question.replace("Image-1: <image>\n\n", "<image>\n")382            history = [383                [item[0].replace("Image-1: <image>\n\n", "<image>\n"), item[1]]384                for item in history385            ]386 387        # Create a new list to store processed sublists388        flattened_list = []389        # Iterate through all but the last sublist in max_input_tile_list and process them390        for sublist in max_input_tile_list[:-1]:391            processed_sublist = [1] * len(392                sublist393            )  # Change each element in the sublist to 1394            flattened_list.extend(395                processed_sublist396            )  # Flatten the processed sublist and add to the new list397        # If max_input_tile_list is not empty, add the last sublist to the new list398        if max_input_tile_list:399            flattened_list.extend(max_input_tile_list[-1])400        max_input_tile_list = flattened_list401        assert len(max_input_tile_list) == len(402            pil_images403        ), "The number of max_input_tile_list and pil_images should be the same."404 405        old_system_message = self.model.system_message406        self.model.system_message = system_message407        image_tiles = []408        transform = build_transform(input_size=self.image_size)409        if len(pil_images) > 0:410            for current_max_input_tiles, pil_image in zip(411                max_input_tile_list, pil_images412            ):413                if self.model.config.dynamic_image_size:414                    tiles = dynamic_preprocess(415                        pil_image,416                        image_size=self.image_size,417                        max_num=current_max_input_tiles,418                        use_thumbnail=self.model.config.use_thumbnail,419                    )420                else:421                    tiles = [pil_image]422                image_tiles += tiles423            pixel_values = [transform(item) for item in image_tiles]424            pixel_values = torch.stack(pixel_values).to(425                self.model.device, dtype=torch.bfloat16426            )427            logger.info(f"Split images to {pixel_values.shape}")428        else:429            pixel_values = None430 431        streamer = TextIteratorStreamer(432            self.tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=10433        )434        generation_config = dict(435            num_beams=1,436            max_new_tokens=max_new_tokens,437            do_sample=do_sample,438            temperature=temperature,439            repetition_penalty=repetition_penalty,440            max_length=self.context_len,441            top_p=top_p,442            streamer=streamer,443        )444        logger.info(f"Generation config: {generation_config}")445        multi_thread_infer(446            self.model,447            self.tokenizer,448            pixel_values,449            question,450            history,451            generation_config,452        )453 454        generated_text = ""455        for new_text in streamer:456            generated_text += new_text457            if generated_text.endswith(self.model.conv_template.sep):458                generated_text = generated_text[: -len(self.model.conv_template.sep)]459            yield json.dumps({"text": generated_text, "error_code": 0}).encode() + b"\0"460        logger.info(461            f"max_input_tile_list: {max_input_tile_list}, history: {history}, "462            f"question: {question}, answer: {generated_text}"463        )464        self.model.system_message = old_system_message465 466    def generate_stream_gate(self, params):467        try:468            for x in self.generate_stream(params):469                yield x470        except ValueError as e:471            print("Caught ValueError:", e)472            traceback.print_exc()473            ret = {474                "text": server_error_msg,475                "error_code": 1,476            }477            yield json.dumps(ret).encode() + b"\0"478        except torch.cuda.CudaError as e:479            traceback.print_exc()480            print("Caught torch.cuda.CudaError:", e)481            ret = {482                "text": server_error_msg,483                "error_code": 1,484            }485            yield json.dumps(ret).encode() + b"\0"486        except Exception as e:487            traceback.print_exc()488            print("Caught Unknown Error", e)489            ret = {490                "text": server_error_msg,491                "error_code": 1,492            }493            yield json.dumps(ret).encode() + b"\0"494 495 496app = FastAPI()497 498 499def release_model_semaphore(fn=None):500    model_semaphore.release()501    if fn is not None:502        fn()503 504 505@app.post("/worker_generate_stream")506async def generate_stream(request: Request):507    global model_semaphore, global_counter508    global_counter += 1509    params = await request.json()510 511    if model_semaphore is None:512        model_semaphore = asyncio.Semaphore(args.limit_model_concurrency)513    await model_semaphore.acquire()514    worker.send_heart_beat()515    generator = worker.generate_stream_gate(params)516    background_tasks = BackgroundTasks()517    background_tasks.add_task(518        partial(release_model_semaphore, fn=worker.send_heart_beat)519    )520    return StreamingResponse(generator, background=background_tasks)521 522 523@app.post("/worker_get_status")524async def get_status(request: Request):525    return worker.get_status()526 527 528if __name__ == "__main__":529    parser = argparse.ArgumentParser()530    parser.add_argument("--host", type=str, default="0.0.0.0")531    parser.add_argument("--port", type=int, default=21002)532    parser.add_argument("--worker-url", type=str, default="http://localhost")533    parser.add_argument("--controller-url", type=str, default="http://localhost:21001")534    parser.add_argument("--model-path", type=str, default="facebook/opt-350m")535    parser.add_argument("--model-name", type=str)536    parser.add_argument("--device", type=str, default="cuda")537    parser.add_argument("--limit-model-concurrency", type=int, default=5)538    parser.add_argument("--stream-interval", type=int, default=1)539    parser.add_argument("--load-8bit", action="store_true")540    args = parser.parse_args()541    logger.info(f"args: {args}")542 543    worker = ModelWorker(544        args.controller_url,545        args.worker_url + f":{args.port}",546        worker_id,547        args.model_path,548        args.model_name,549        args.load_8bit,550        args.device,551    )552    uvicorn.run(app, host=args.host, port=args.port, log_level="info")553