CoolFace
Apppublic

q-future/Co-Instruct

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
29likes
model_worker.py227 linesDownload Raw Back to root
1"""2A model worker executes the model.3"""4import argparse5import asyncio6import json7import time8import threading9import uuid10 11import requests12import torch13from functools import partial14 15from mplug_owl2.constants import WORKER_HEART_BEAT_INTERVAL16from mplug_owl2.utils import (build_logger, server_error_msg,17    pretty_print_semaphore)18from mplug_owl2.model.builder import load_pretrained_model19from mplug_owl2.mm_utils import process_images, load_image_from_base64, tokenizer_image_token, KeywordsStoppingCriteria20from mplug_owl2.constants import IMAGE_TOKEN_INDEX, DEFAULT_IMAGE_TOKEN21from transformers import TextIteratorStreamer22from threading import Thread23 24GB = 1 << 3025 26worker_id = str(uuid.uuid4())[:6]27logger = build_logger("model_worker", f"model_worker_{worker_id}.log")28 29class ModelWorker:30    def __init__(self, model_path, model_base, model_name, load_8bit, load_4bit, device):31        self.worker_id = worker_id32        if model_path.endswith("/"):33            model_path = model_path[:-1]34        if model_name is None:35            model_paths = model_path.split("/")36            if model_paths[-1].startswith('checkpoint-'):37                self.model_name = model_paths[-2] + "_" + model_paths[-1]38            else:39                self.model_name = model_paths[-1]40        else:41            self.model_name = model_name42 43        self.device = device44        logger.info(f"Loading the model {self.model_name} on worker {worker_id} ...")45        self.tokenizer, self.model, self.image_processor, self.context_len = load_pretrained_model(46            model_path, model_base, self.model_name, load_8bit, load_4bit, device=self.device)47        self.is_multimodal = True48        49    @torch.inference_mode()50    def predict_stream(self, params):51        tokenizer, model, image_processor = self.tokenizer, self.model, self.image_processor52 53        prompt = params["prompt"] + "The quality of the image is"54        ori_prompt = prompt55        images = params.get("images", None)56        num_image_tokens = 057        if images is not None and len(images) > 0 and self.is_multimodal:58            if len(images) > 0:59                if len(images) != prompt.count(DEFAULT_IMAGE_TOKEN):60                    raise ValueError("Number of images does not match number of <|image|> tokens in prompt")61 62                images = [load_image_from_base64(image) for image in images]63                images = process_images(images, image_processor, model.config)64 65                if type(images) is list:66                    images = [image.to(self.model.device, dtype=torch.float16) for image in images]67                else:68                    images = images.to(self.model.device, dtype=torch.float16)69 70                replace_token = DEFAULT_IMAGE_TOKEN71                prompt = prompt.replace(DEFAULT_IMAGE_TOKEN, replace_token)72 73                num_image_tokens = prompt.count(replace_token) * (model.get_model().visual_abstractor.config.num_learnable_queries + 1)74            else:75                images = None76            image_args = {"images": images}77        else:78            images = None79            image_args = {}80            81        input_ids = tokenizer_image_token(prompt, tokenizer, IMAGE_TOKEN_INDEX, return_tensors='pt').unsqueeze(0).to(self.device)82        83        logits = model.forward(84            input_ids=input_ids,85            use_cache=True,86            **image_args).logits[0,-1]87        88        print(logits.shape)89        90        softmax_logits = torch.softmax(logits[[1781,6588,6460]], 0)91        92        print(tokenizer(["good", "average", "poor"]))93        fake_streamer = []94        for id_, word in enumerate(["good", "average", "poor"]):95            stream_ = f"Probability of {word} quality: {softmax_logits[id_].item():.4f};\n"96            fake_streamer.append(stream_)97        98        quality_score = 0.5 * softmax_logits[1] + softmax_logits[0]99        stream_ = f"Quality score: {quality_score:.4f} (range [0,1])."100        fake_streamer.append(stream_)101        102        generated_text = ori_prompt.replace("The quality of the image is", "")103        for new_text in fake_streamer:104            generated_text += new_text105            yield json.dumps({"text": generated_text, "error_code": 0}).encode()106    107    @torch.inference_mode()108    def generate_stream(self, params):109        tokenizer, model, image_processor = self.tokenizer, self.model, self.image_processor110 111        prompt = params["prompt"]112        ori_prompt = prompt113        images = params.get("images", None)114        num_image_tokens = 0115        if images is not None and len(images) > 0 and self.is_multimodal:116            if len(images) > 0:117                if len(images) != prompt.count(DEFAULT_IMAGE_TOKEN):118                    raise ValueError("Number of images does not match number of <|image|> tokens in prompt")119 120                images = [load_image_from_base64(image) for image in images]121                images = process_images(images, image_processor, model.config)122 123                if type(images) is list:124                    images = [image.to(self.model.device, dtype=torch.float16) for image in images]125                else:126                    images = images.to(self.model.device, dtype=torch.float16)127 128                replace_token = DEFAULT_IMAGE_TOKEN129                prompt = prompt.replace(DEFAULT_IMAGE_TOKEN, replace_token)130 131                num_image_tokens = prompt.count(replace_token) * (model.get_model().visual_abstractor.config.num_learnable_queries + 1)132            else:133                images = None134            image_args = {"images": images}135        else:136            images = None137            image_args = {}138 139        temperature = float(params.get("temperature", 1.0))140        top_p = float(params.get("top_p", 1.0))141        max_context_length = getattr(model.config, 'max_position_embeddings', 4096)142        max_new_tokens = min(int(params.get("max_new_tokens", 256)), 1024)143        stop_str = params.get("stop", None)144        do_sample = True if temperature > 0.001 else False145 146        input_ids = tokenizer_image_token(prompt, tokenizer, IMAGE_TOKEN_INDEX, return_tensors='pt').unsqueeze(0).to(self.device)147        keywords = [stop_str]148        stopping_criteria = KeywordsStoppingCriteria(keywords, tokenizer, input_ids)149        streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=15)150 151        max_new_tokens = min(max_new_tokens, max_context_length - input_ids.shape[-1] - num_image_tokens)152 153        if max_new_tokens < 1:154            yield json.dumps({"text": ori_prompt + "Exceeds max token length. Please start a new conversation, thanks.", "error_code": 0}).encode() + b"\0"155            return156 157        thread = Thread(target=model.generate, kwargs=dict(158            inputs=input_ids,159            do_sample=do_sample,160            temperature=temperature,161            top_p=top_p,162            max_new_tokens=max_new_tokens,163            streamer=streamer,164            stopping_criteria=[stopping_criteria],165            use_cache=True,166            **image_args167        ))168        thread.start()169 170        generated_text = ori_prompt171        for new_text in streamer:172            generated_text += new_text173            if generated_text.endswith(stop_str):174                generated_text = generated_text[:-len(stop_str)]175            yield json.dumps({"text": generated_text, "error_code": 0}).encode()176            177    def predict_stream_gate(self, params):178        try:179            for x in self.predict_stream(params):180                yield x181        except ValueError as e:182            print("Caught ValueError:", e)183            ret = {184                "text": server_error_msg,185                "error_code": 1,186            }187            yield json.dumps(ret).encode() 188        except torch.cuda.CudaError as e:189            print("Caught torch.cuda.CudaError:", e)190            ret = {191                "text": server_error_msg,192                "error_code": 1,193            }194            yield json.dumps(ret).encode()195        except Exception as e:196            print("Caught Unknown Error", e)197            ret = {198                "text": server_error_msg,199                "error_code": 1,200            }201            yield json.dumps(ret).encode()202 203    def generate_stream_gate(self, params):204        try:205            for x in self.generate_stream(params):206                yield x207        except ValueError as e:208            print("Caught ValueError:", e)209            ret = {210                "text": server_error_msg,211                "error_code": 1,212            }213            yield json.dumps(ret).encode() 214        except torch.cuda.CudaError as e:215            print("Caught torch.cuda.CudaError:", e)216            ret = {217                "text": server_error_msg,218                "error_code": 1,219            }220            yield json.dumps(ret).encode()221        except Exception as e:222            print("Caught Unknown Error", e)223            ret = {224                "text": server_error_msg,225                "error_code": 1,226            }227            yield json.dumps(ret).encode()