globc/LLaVA
0
1import gradio as gr2from llava.mm_utils import get_model_name_from_path3from llava.model.builder import load_pretrained_model4from llava.eval.run_llava import eval_model5 6model_path = "liuhaotian/llava-v1.5-7b"7 8model_name = get_model_name_from_path(model_path)9tokenizer, model, image_processor, context_len = load_pretrained_model(model_path, None, model_name, load_4bit=True)10 11def predict(input_img, prompt):12 13 args = type('Args', (), {14 "model_path": model_path,15 "model_base": None,16 "model_name": model_name,17 "query": prompt,18 "conv_mode": None,19 "image_file": input_img,20 "sep": ",",21 "temperature": 0.2,22 "top_p": None,23 "num_beams": 1,24 "max_new_tokens": 51225 })()26 27 return eval_model(args, tokenizer, model, image_processor, context_len)28 29gradio_app = gr.Interface(30 fn=predict,31 inputs=[gr.Image(type="filepath"), "text"],32 outputs="text"33)34 35gradio_app.launch()