ybelkada commited on
Commit
4b0ba99
Β·
verified Β·
1 Parent(s): d222751

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +146 -0
app.py ADDED
@@ -0,0 +1,146 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+ from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
4
+ import gradio as gr
5
+ from threading import Thread
6
+
7
+ MODEL = "tiiuae/Falcon-EB-3B-Instruct"
8
+
9
+ TITLE = "<h1><center>Falcon-E-3B-Instruct playground</center></h1>"
10
+ SUB_TITLE = """<center>This interface has been created for quick validation purposes, do not use it for production.</center>"""
11
+
12
+ CSS = """
13
+ .duplicate-button {
14
+ margin: auto !important;
15
+ color: white !important;
16
+ background: black !important;
17
+ border-radius: 100vh !important;
18
+ }
19
+ h3 {
20
+ text-align: center;
21
+ }
22
+ """
23
+
24
+ END_MESSAGE = """
25
+ \n
26
+ **The conversation has reached to its end, please press "Clear" to restart a new conversation**
27
+ """
28
+
29
+ device = "cuda" # for GPU usage or "cpu" for CPU usage
30
+
31
+ tokenizer = AutoTokenizer.from_pretrained(MODEL)
32
+ model = AutoModelForCausalLM.from_pretrained(
33
+ MODEL,
34
+ torch_dtype=torch.bfloat16,
35
+ ).to(device)
36
+
37
+ model = torch.compile(model)
38
+
39
+ def stream_chat(
40
+ message: str,
41
+ history: list,
42
+ temperature: float = 0.3,
43
+ max_new_tokens: int = 128,
44
+ top_p: float = 1.0,
45
+ top_k: int = 20,
46
+ penalty: float = 1.2,
47
+ ):
48
+ print(f'message: {message}')
49
+ print(f'history: {history}')
50
+
51
+ conversation = []
52
+ for prompt, answer in history:
53
+ conversation.extend([
54
+ {"role": "user", "content": prompt},
55
+ {"role": "assistant", "content": answer},
56
+ ])
57
+
58
+
59
+ conversation.append({"role": "user", "content": message})
60
+ input_text = tokenizer.apply_chat_template(conversation, tokenize=False, add_generation_prompt = True)
61
+
62
+ inputs = tokenizer.encode(input_text, return_tensors="pt").to(device)
63
+ streamer = TextIteratorStreamer(tokenizer, timeout=60.0, skip_prompt=True, skip_special_tokens=True)
64
+
65
+ generate_kwargs = dict(
66
+ input_ids=inputs,
67
+ max_new_tokens = max_new_tokens,
68
+ do_sample = False if temperature == 0 else True,
69
+ top_p = top_p,
70
+ top_k = top_k,
71
+ temperature = temperature,
72
+ streamer=streamer,
73
+ pad_token_id = 10,
74
+ )
75
+
76
+ with torch.no_grad():
77
+ thread = Thread(target=model.generate, kwargs=generate_kwargs)
78
+ thread.start()
79
+
80
+ buffer = ""
81
+ for new_text in streamer:
82
+ buffer += new_text
83
+ yield buffer
84
+
85
+
86
+ print(f'response: {buffer}')
87
+
88
+ chatbot = gr.Chatbot(height=600)
89
+
90
+ with gr.Blocks(css=CSS, theme="soft") as demo:
91
+ gr.HTML(TITLE)
92
+ gr.HTML(SUB_TITLE)
93
+ gr.DuplicateButton(value="Duplicate Space for private use", elem_classes="duplicate-button")
94
+ gr.ChatInterface(
95
+ fn=stream_chat,
96
+ chatbot=chatbot,
97
+ fill_height=True,
98
+ additional_inputs_accordion=gr.Accordion(label="βš™οΈ Parameters", open=False, render=False),
99
+ additional_inputs=[
100
+ gr.Slider(
101
+ minimum=0,
102
+ maximum=1,
103
+ step=0.1,
104
+ value=0.3,
105
+ label="Temperature",
106
+ render=False,
107
+ ),
108
+ gr.Slider(
109
+ minimum=128,
110
+ maximum=4096,
111
+ step=1,
112
+ value=128,
113
+ label="Max new tokens",
114
+ render=False,
115
+ ),
116
+ gr.Slider(
117
+ minimum=0.0,
118
+ maximum=1.0,
119
+ step=0.1,
120
+ value=1.0,
121
+ label="top_p",
122
+ render=False,
123
+ ),
124
+ gr.Slider(
125
+ minimum=1,
126
+ maximum=20,
127
+ step=1,
128
+ value=20,
129
+ label="top_k",
130
+ render=False,
131
+ ),
132
+ gr.Slider(
133
+ minimum=0.0,
134
+ maximum=2.0,
135
+ step=0.1,
136
+ value=1.2,
137
+ label="Repetition penalty",
138
+ render=False,
139
+ ),
140
+ ],
141
+ cache_examples=False,
142
+ )
143
+
144
+
145
+ if __name__ == "__main__":
146
+ demo.launch()