Sapiensia/LLaVA
1
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 