-
Notifications
You must be signed in to change notification settings - Fork 1k
Expand file tree
/
Copy pathgradio_openchatkit.py
More file actions
90 lines (70 loc) · 2.82 KB
/
Copy pathgradio_openchatkit.py
File metadata and controls
90 lines (70 loc) · 2.82 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
#!/usr/bin/env python
# -*- coding: utf-8 -*-
# @Desc : openchat kit gradio application
"""
Run:
# under OpenChatKit/inference from https://github.com/togethercomputer/OpenChatKit
CUDA_VISIBLE_DEVICES=2,3 python3 gradio_openchatkit.py
Warn:
the bigger max_new_tokens the more cuda mem, so be careful
"""
import os
import sys
CUR_DIR = os.path.abspath(os.path.dirname(__file__))
MODEL_PATH = os.path.join(CUR_DIR, "../../GPT-NeoXT-Chat-Base-20B/")
sys.path.append(CUR_DIR)
from loguru import logger
import gradio as gr
import argparse
from transformers import AutoTokenizer, AutoModelForCausalLM
from bot import ChatModel
class ConvChat(object):
"""
Conversation Chat
"""
def __init__(self,
model_name: str,
max_new_tokens: int = 256,
sample: bool = False,
temperature: int = 0.6,
top_k: int = 40):
self.max_new_tokens = max_new_tokens
self.sample = sample
self.temperature = temperature
self.top_k = top_k
logger.info("Start to init Chat Model")
self.chat_model = ChatModel(model_name=model_name, gpu_id=0)
logger.info("Initialized Chat Model")
def run_text(self, input_text: gr.Textbox, state: gr.State):
response = self.chat_model.do_inference(
prompt=input_text,
max_new_tokens=self.max_new_tokens,
do_sample=self.sample,
temperature=self.temperature,
top_k=self.top_k
)
state = state + [(input_text, response)]
return state, state
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_name", type=str, default=MODEL_PATH, help="model huggingface repo name or local path")
parser.add_argument("--server_port", type=int, default=7800, help="gradio server port")
args = parser.parse_args()
conv_chat = ConvChat(model_name=args.model_name)
with gr.Blocks(css="OpenChatKit .overflow-y-auto{height:500px}") as gr_chat:
chatbot = gr.Chatbot(elem_id="chatbot", label="OpenChatKit")
state = gr.State([])
with gr.Row():
with gr.Column(scale=0.8):
input_text = gr.Textbox(show_label=False,
placeholder="Enter your question").style(container=False)
with gr.Column(scale=0.2, min_width=0):
clear_btn = gr.Button("Clear")
input_text.submit(conv_chat.run_text, [input_text, state], [chatbot, state])
input_text.submit(lambda: "", None, input_text)
clear_btn.click(lambda: [], None, chatbot)
clear_btn.click(lambda: [], None, state)
gr_chat.launch(
server_name="0.0.0.0",
server_port=args.server_port
)