OpenGVLab/InternVL
510
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 