CoolFace
Apppublic

Sapiensia/LLaVA

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
app.py627 linesDownload Raw Back to root
1import argparse2import datetime3import hashlib4import json5import os6import subprocess7import sys8import time9 10import gradio as gr11import requests12 13from llava.constants import LOGDIR14from llava.conversation import SeparatorStyle, conv_templates, default_conversation15from llava.utils import (16    build_logger,17    moderation_msg,18    server_error_msg,19    violates_moderation,20)21 22logger = build_logger("gradio_web_server", "gradio_web_server.log")23 24headers = {"User-Agent": "LLaVA Client"}25 26no_change_btn = gr.Button.update()27enable_btn = gr.Button.update(interactive=True)28disable_btn = gr.Button.update(interactive=False)29 30priority = {31    "vicuna-13b": "aaaaaaa",32    "koala-13b": "aaaaaab",33}34 35 36def get_conv_log_filename():37    t = datetime.datetime.now()38    name = os.path.join(LOGDIR, f"{t.year}-{t.month:02d}-{t.day:02d}-conv.json")39    return name40 41 42def get_model_list():43    ret = requests.post(args.controller_url + "/refresh_all_workers")44    assert ret.status_code == 20045    ret = requests.post(args.controller_url + "/list_models")46    models = ret.json()["models"]47    models.sort(key=lambda x: priority.get(x, x))48    logger.info(f"Models: {models}")49    return models50 51 52get_window_url_params = """53function() {54    const params = new URLSearchParams(window.location.search);55    url_params = Object.fromEntries(params);56    console.log(url_params);57    return url_params;58    }59"""60 61 62def load_demo(url_params, request: gr.Request):63    logger.info(f"load_demo. ip: {request.client.host}. params: {url_params}")64 65    dropdown_update = gr.Dropdown.update(visible=True)66    if "model" in url_params:67        model = url_params["model"]68        if model in models:69            dropdown_update = gr.Dropdown.update(value=model, visible=True)70 71    state = default_conversation.copy()72    return state, dropdown_update73 74 75def load_demo_refresh_model_list(request: gr.Request):76    logger.info(f"load_demo. ip: {request.client.host}")77    models = get_model_list()78    state = default_conversation.copy()79 80    models_downloaded = True if models else False81 82    model_dropdown_kwargs = {83        "choices": [],84        "value": "Downloading the models...",85        "interactive": models_downloaded,86    }87 88    if models_downloaded:89        model_dropdown_kwargs["choices"] = models90        model_dropdown_kwargs["value"] = models[0]91 92    models_dropdown_update = gr.Dropdown.update(**model_dropdown_kwargs)93 94    send_button_update = gr.Button.update(95        interactive=models_downloaded,96    )97 98    return state, models_dropdown_update, send_button_update99 100 101def vote_last_response(state, vote_type, model_selector, request: gr.Request):102    with open(get_conv_log_filename(), "a") as fout:103        data = {104            "tstamp": round(time.time(), 4),105            "type": vote_type,106            "model": model_selector,107            "state": state.dict(),108            "ip": request.client.host,109        }110        fout.write(json.dumps(data) + "\n")111 112 113def upvote_last_response(state, model_selector, request: gr.Request):114    logger.info(f"upvote. ip: {request.client.host}")115    vote_last_response(state, "upvote", model_selector, request)116    return ("",) + (disable_btn,) * 3117 118 119def downvote_last_response(state, model_selector, request: gr.Request):120    logger.info(f"downvote. ip: {request.client.host}")121    vote_last_response(state, "downvote", model_selector, request)122    return ("",) + (disable_btn,) * 3123 124 125def flag_last_response(state, model_selector, request: gr.Request):126    logger.info(f"flag. ip: {request.client.host}")127    vote_last_response(state, "flag", model_selector, request)128    return ("",) + (disable_btn,) * 3129 130 131def regenerate(state, image_process_mode, request: gr.Request):132    logger.info(f"regenerate. ip: {request.client.host}")133    state.messages[-1][-1] = None134    prev_human_msg = state.messages[-2]135    if type(prev_human_msg[1]) in (tuple, list):136        prev_human_msg[1] = (*prev_human_msg[1][:2], image_process_mode)137    state.skip_next = False138    return (state, state.to_gradio_chatbot(), "", None) + (disable_btn,) * 5139 140 141def clear_history(request: gr.Request):142    logger.info(f"clear_history. ip: {request.client.host}")143    state = default_conversation.copy()144    return (state, state.to_gradio_chatbot(), "", None) + (disable_btn,) * 5145 146 147def add_text(state, text, image, image_process_mode, request: gr.Request):148    logger.info(f"add_text. ip: {request.client.host}. len: {len(text)}")149    if len(text) <= 0 and image is None:150        state.skip_next = True151        return (state, state.to_gradio_chatbot(), "", None) + (no_change_btn,) * 5152    if args.moderate:153        flagged = violates_moderation(text)154        if flagged:155            state.skip_next = True156            return (state, state.to_gradio_chatbot(), moderation_msg, None) + (157                no_change_btn,158            ) * 5159 160    text = text[:1536]  # Hard cut-off161    if image is not None:162        text = text[:1200]  # Hard cut-off for images163        if "<image>" not in text:164            # text = '<Image><image></Image>' + text165            text = text + "\n<image>"166        text = (text, image, image_process_mode)167        if len(state.get_images(return_pil=True)) > 0:168            state = default_conversation.copy()169    state.append_message(state.roles[0], text)170    state.append_message(state.roles[1], None)171    state.skip_next = False172    return (state, state.to_gradio_chatbot(), "", None) + (disable_btn,) * 5173 174 175def http_bot(176    state, model_selector, temperature, top_p, max_new_tokens, request: gr.Request177):178    logger.info(f"http_bot. ip: {request.client.host}")179    start_tstamp = time.time()180    model_name = model_selector181 182    if state.skip_next:183        # This generate call is skipped due to invalid inputs184        yield (state, state.to_gradio_chatbot()) + (no_change_btn,) * 5185        return186 187    if len(state.messages) == state.offset + 2:188        # First round of conversation189        if "llava" in model_name.lower():190            if "llama-2" in model_name.lower():191                template_name = "llava_llama_2"192            elif "v1" in model_name.lower():193                if "mmtag" in model_name.lower():194                    template_name = "v1_mmtag"195                elif (196                    "plain" in model_name.lower()197                    and "finetune" not in model_name.lower()198                ):199                    template_name = "v1_mmtag"200                else:201                    template_name = "llava_v1"202            elif "mpt" in model_name.lower():203                template_name = "mpt"204            else:205                if "mmtag" in model_name.lower():206                    template_name = "v0_mmtag"207                elif (208                    "plain" in model_name.lower()209                    and "finetune" not in model_name.lower()210                ):211                    template_name = "v0_mmtag"212                else:213                    template_name = "llava_v0"214        elif "mpt" in model_name:215            template_name = "mpt_text"216        elif "llama-2" in model_name:217            template_name = "llama_2"218        else:219            template_name = "vicuna_v1"220        new_state = conv_templates[template_name].copy()221        new_state.append_message(new_state.roles[0], state.messages[-2][1])222        new_state.append_message(new_state.roles[1], None)223        state = new_state224 225    # Query worker address226    controller_url = args.controller_url227    ret = requests.post(228        controller_url + "/get_worker_address", json={"model": model_name}229    )230    worker_addr = ret.json()["address"]231    logger.info(f"model_name: {model_name}, worker_addr: {worker_addr}")232 233    # No available worker234    if worker_addr == "":235        state.messages[-1][-1] = server_error_msg236        yield (237            state,238            state.to_gradio_chatbot(),239            disable_btn,240            disable_btn,241            disable_btn,242            enable_btn,243            enable_btn,244        )245        return246 247    # Construct prompt248    prompt = state.get_prompt()249 250    all_images = state.get_images(return_pil=True)251    all_image_hash = [hashlib.md5(image.tobytes()).hexdigest() for image in all_images]252    for image, hash in zip(all_images, all_image_hash):253        t = datetime.datetime.now()254        filename = os.path.join(255            LOGDIR, "serve_images", f"{t.year}-{t.month:02d}-{t.day:02d}", f"{hash}.jpg"256        )257        if not os.path.isfile(filename):258            os.makedirs(os.path.dirname(filename), exist_ok=True)259            image.save(filename)260 261    # Make requests262    pload = {263        "model": model_name,264        "prompt": prompt,265        "temperature": float(temperature),266        "top_p": float(top_p),267        "max_new_tokens": min(int(max_new_tokens), 1536),268        "stop": state.sep269        if state.sep_style in [SeparatorStyle.SINGLE, SeparatorStyle.MPT]270        else state.sep2,271        "images": f"List of {len(state.get_images())} images: {all_image_hash}",272    }273    logger.info(f"==== request ====\n{pload}")274 275    pload["images"] = state.get_images()276 277    state.messages[-1][-1] = "โ–Œ"278    yield (state, state.to_gradio_chatbot()) + (disable_btn,) * 5279 280    try:281        # Stream output282        response = requests.post(283            worker_addr + "/worker_generate_stream",284            headers=headers,285            json=pload,286            stream=True,287            timeout=10,288        )289        for chunk in response.iter_lines(decode_unicode=False, delimiter=b"\0"):290            if chunk:291                data = json.loads(chunk.decode())292                if data["error_code"] == 0:293                    output = data["text"][len(prompt) :].strip()294                    state.messages[-1][-1] = output + "โ–Œ"295                    yield (state, state.to_gradio_chatbot()) + (disable_btn,) * 5296                else:297                    output = data["text"] + f" (error_code: {data['error_code']})"298                    state.messages[-1][-1] = output299                    yield (state, state.to_gradio_chatbot()) + (300                        disable_btn,301                        disable_btn,302                        disable_btn,303                        enable_btn,304                        enable_btn,305                    )306                    return307                time.sleep(0.03)308    except requests.exceptions.RequestException as e:309        state.messages[-1][-1] = server_error_msg310        yield (state, state.to_gradio_chatbot()) + (311            disable_btn,312            disable_btn,313            disable_btn,314            enable_btn,315            enable_btn,316        )317        return318 319    state.messages[-1][-1] = state.messages[-1][-1][:-1]320    yield (state, state.to_gradio_chatbot()) + (enable_btn,) * 5321 322    finish_tstamp = time.time()323    logger.info(f"{output}")324 325    with open(get_conv_log_filename(), "a") as fout:326        data = {327            "tstamp": round(finish_tstamp, 4),328            "type": "chat",329            "model": model_name,330            "start": round(start_tstamp, 4),331            "finish": round(start_tstamp, 4),332            "state": state.dict(),333            "images": all_image_hash,334            "ip": request.client.host,335        }336        fout.write(json.dumps(data) + "\n")337 338 339title_markdown = """340# ๐ŸŒ‹ LLaVA: Large Language and Vision Assistant341[[Project Page]](https://llava-vl.github.io) [[Paper]](https://arxiv.org/abs/2304.08485) [[Code]](https://github.com/haotian-liu/LLaVA) [[Model]](https://github.com/haotian-liu/LLaVA/blob/main/docs/MODEL_ZOO.md)342 343ONLY WORKS WITH GPU!344 345You can load the model with 4-bit or 8-bit quantization to make it fit in smaller hardwares. Setting the environment variable `bits` to control the quantization.346*Note: 8-bit seems to be slower than both 4-bit/16-bit. Although it has enough VRAM to support 8-bit, until we figure out the inference speed issue, we recommend 4-bit for A10G for the best efficiency.*347 348Recommended configurations:349| Hardware          | T4-Small (16G)  | A10G-Small (24G) | A100-Large (40G) |350|-------------------|-----------------|------------------|------------------|351| **Bits**          | 4 (default)     | 4                | 16               |352 353"""354 355tos_markdown = """356### Terms of use357By using this service, users are required to agree to the following terms:358The service is a research preview intended for non-commercial use only. It only provides limited safety measures and may generate offensive content. It must not be used for any illegal, harmful, violent, racist, or sexual purposes. The service may collect user dialogue data for future research.359Please click the "Flag" button if you get any inappropriate answer! We will collect those to keep improving our moderator.360For an optimal experience, please use desktop computers for this demo, as mobile devices may compromise its quality.361"""362 363 364learn_more_markdown = """365### License366The service is a research preview intended for non-commercial use only, subject to the model [License](https://github.com/facebookresearch/llama/blob/main/MODEL_CARD.md) of LLaMA, [Terms of Use](https://openai.com/policies/terms-of-use) of the data generated by OpenAI, and [Privacy Practices](https://chrome.google.com/webstore/detail/sharegpt-share-your-chatg/daiacboceoaocpibfodeljbdfacokfjb) of ShareGPT. Please contact us if you find any potential violation.367"""368 369block_css = """370 371#buttons button {372    min-width: min(120px,100%);373}374 375"""376 377 378def build_demo(embed_mode):379    models = get_model_list()380 381    textbox = gr.Textbox(382        show_label=False, placeholder="Enter text and press ENTER", container=False383    )384    with gr.Blocks(title="LLaVA", theme=gr.themes.Default(), css=block_css) as demo:385        state = gr.State(default_conversation.copy())386 387        if not embed_mode:388            gr.Markdown(title_markdown)389 390        with gr.Row():391            with gr.Column(scale=3):392                with gr.Row(elem_id="model_selector_row"):393                    model_selector = gr.Dropdown(394                        choices=models,395                        value=models[0] if models else "Downloading the models...",396                        interactive=True if models else False,397                        show_label=False,398                        container=False,399                    )400 401                imagebox = gr.Image(type="pil")402                image_process_mode = gr.Radio(403                    ["Crop", "Resize", "Pad", "Default"],404                    value="Default",405                    label="Preprocess for non-square image",406                    visible=False,407                )408 409                cur_dir = os.path.dirname(os.path.abspath(__file__))410                gr.Examples(411                    examples=[412                        [413                            f"{cur_dir}/examples/extreme_ironing.jpg",414                            "What is unusual about this image?",415                        ],416                        [417                            f"{cur_dir}/examples/waterview.jpg",418                            "What are the things I should be cautious about when I visit here?",419                        ],420                    ],421                    inputs=[imagebox, textbox],422                )423 424                with gr.Accordion("Parameters", open=False) as parameter_row:425                    temperature = gr.Slider(426                        minimum=0.0,427                        maximum=1.0,428                        value=0.2,429                        step=0.1,430                        interactive=True,431                        label="Temperature",432                    )433                    top_p = gr.Slider(434                        minimum=0.0,435                        maximum=1.0,436                        value=0.7,437                        step=0.1,438                        interactive=True,439                        label="Top P",440                    )441                    max_output_tokens = gr.Slider(442                        minimum=0,443                        maximum=1024,444                        value=512,445                        step=64,446                        interactive=True,447                        label="Max output tokens",448                    )449 450            with gr.Column(scale=8):451                chatbot = gr.Chatbot(452                    elem_id="chatbot", label="LLaVA Chatbot", height=550453                )454                with gr.Row():455                    with gr.Column(scale=8):456                        textbox.render()457                    with gr.Column(scale=1, min_width=50):458                        submit_btn = gr.Button(459                            value="Send", variant="primary", interactive=False460                        )461                with gr.Row(elem_id="buttons") as button_row:462                    upvote_btn = gr.Button(value="๐Ÿ‘  Upvote", interactive=False)463                    downvote_btn = gr.Button(value="๐Ÿ‘Ž  Downvote", interactive=False)464                    flag_btn = gr.Button(value="โš ๏ธ  Flag", interactive=False)465                    # stop_btn = gr.Button(value="โน๏ธ  Stop Generation", interactive=False)466                    regenerate_btn = gr.Button(value="๐Ÿ”„  Regenerate", interactive=False)467                    clear_btn = gr.Button(value="๐Ÿ—‘๏ธ  Clear history", interactive=False)468 469        if not embed_mode:470            gr.Markdown(tos_markdown)471            gr.Markdown(learn_more_markdown)472        url_params = gr.JSON(visible=False)473 474        # Register listeners475        btn_list = [upvote_btn, downvote_btn, flag_btn, regenerate_btn, clear_btn]476        upvote_btn.click(477            upvote_last_response,478            [state, model_selector],479            [textbox, upvote_btn, downvote_btn, flag_btn],480        )481        downvote_btn.click(482            downvote_last_response,483            [state, model_selector],484            [textbox, upvote_btn, downvote_btn, flag_btn],485        )486        flag_btn.click(487            flag_last_response,488            [state, model_selector],489            [textbox, upvote_btn, downvote_btn, flag_btn],490        )491        regenerate_btn.click(492            regenerate,493            [state, image_process_mode],494            [state, chatbot, textbox, imagebox] + btn_list,495        ).then(496            http_bot,497            [state, model_selector, temperature, top_p, max_output_tokens],498            [state, chatbot] + btn_list,499        )500        clear_btn.click(501            clear_history, None, [state, chatbot, textbox, imagebox] + btn_list502        )503 504        textbox.submit(505            add_text,506            [state, textbox, imagebox, image_process_mode],507            [state, chatbot, textbox, imagebox] + btn_list,508        ).then(509            http_bot,510            [state, model_selector, temperature, top_p, max_output_tokens],511            [state, chatbot] + btn_list,512        )513        submit_btn.click(514            add_text,515            [state, textbox, imagebox, image_process_mode],516            [state, chatbot, textbox, imagebox] + btn_list,517        ).then(518            http_bot,519            [state, model_selector, temperature, top_p, max_output_tokens],520            [state, chatbot] + btn_list,521        )522 523        if args.model_list_mode == "once":524            demo.load(525                load_demo,526                [url_params],527                [state, model_selector],528                _js=get_window_url_params,529            )530        elif args.model_list_mode == "reload":531            demo.load(532                load_demo_refresh_model_list, None, [state, model_selector, submit_btn]533            )534        else:535            raise ValueError(f"Unknown model list mode: {args.model_list_mode}")536 537    return demo538 539 540def start_controller():541    logger.info("Starting the controller")542    controller_command = [543        "python",544        "-m",545        "llava.serve.controller",546        "--host",547        "0.0.0.0",548        "--port",549        "10000",550    ]551    return subprocess.Popen(controller_command)552 553 554def start_worker(model_path: str, bits=16):555    logger.info(f"Starting the model worker for the model {model_path}")556    model_name = model_path.strip("/").split("/")[-1]557    assert bits in [4, 8, 16], "It can be only loaded with 16-bit, 8-bit, and 4-bit."558    if bits != 16:559        model_name += f"-{bits}bit"560    worker_command = [561        "python",562        "-m",563        "llava.serve.model_worker",564        "--host",565        "0.0.0.0",566        "--controller",567        "http://localhost:10000",568        "--model-path",569        model_path,570        "--model-name",571        model_name,572    ]573    if bits != 16:574        worker_command += [f"--load-{bits}bit"]575    return subprocess.Popen(worker_command)576 577 578def get_args():579    parser = argparse.ArgumentParser()580    parser.add_argument("--host", type=str, default="0.0.0.0")581    parser.add_argument("--port", type=int)582    parser.add_argument("--controller-url", type=str, default="http://localhost:10000")583    parser.add_argument("--concurrency-count", type=int, default=8)584    parser.add_argument(585        "--model-list-mode", type=str, default="reload", choices=["once", "reload"]586    )587    parser.add_argument("--share", action="store_true")588    parser.add_argument("--moderate", action="store_true")589    parser.add_argument("--embed", action="store_true")590 591    args = parser.parse_args()592 593    return args594 595 596def start_demo(args):597    demo = build_demo(args.embed)598    demo.queue(599        concurrency_count=args.concurrency_count, status_update_rate=10, api_open=False600    ).launch(server_name=args.host, server_port=args.port, share=args.share)601 602 603if __name__ == "__main__":604    args = get_args()605    logger.info(f"args: {args}")606 607    model_path = "liuhaotian/llava-v1.5-13b"608    bits = int(os.getenv("bits", 8))609 610    controller_proc = start_controller()611    worker_proc = start_worker(model_path, bits=bits)612 613    # Wait for worker and controller to start614    time.sleep(10)615 616    exit_status = 0617    try:618        start_demo(args)619    except Exception as e:620        print(e)621        exit_status = 1622    finally:623        worker_proc.kill()624        controller_proc.kill()625 626        sys.exit(exit_status)627