devoppro commited on
Commit
27f7554
·
verified ·
1 Parent(s): 5c5f48d

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +314 -66
app.py CHANGED
@@ -1,98 +1,346 @@
1
- import spaces
2
- import gradio as gr
3
  import torch
4
- from transformers import AutoModelForCausalLM, AutoTokenizer
 
 
5
 
6
  MODEL_ID = "devoppro/FastLLM"
7
- TOKENIZER_ID = "Qwen/Qwen2.5-0.5B" # FastLLM reuses the Qwen2.5 BPE vocab
8
 
9
- tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_ID)
10
- model = AutoModelForCausalLM.from_pretrained(
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
11
  MODEL_ID,
12
- torch_dtype=torch.float16,
13
  trust_remote_code=True,
14
  )
15
- model.to("cuda")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
16
  model.eval()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
17
 
 
 
 
 
18
 
19
- @spaces.GPU(duration=30)
20
- def generate(prompt, max_new_tokens, temperature, top_k):
21
- input_ids = tokenizer.encode(prompt, return_tensors="pt").to("cuda")
22
- max_new_tokens = int(max_new_tokens)
23
- temperature = float(temperature)
24
- top_k = int(top_k)
25
 
26
- with torch.no_grad():
27
- for _ in range(max_new_tokens):
28
- outputs = model(input_ids)
29
 
30
- # Cast to fp32 for the sampling math — fp16 logits from an
31
- # early/undertrained checkpoint can overflow fp16's range,
32
- # which turns softmax output into inf/nan and crashes
33
- # multinomial with a CUDA device-side assert.
34
- logits = outputs["logits"][:, -1, :].float() / max(temperature, 1e-5)
35
 
36
- if top_k > 0:
37
- v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
38
- logits[logits < v[:, [-1]]] = -float("Inf")
39
 
40
- probs = torch.softmax(logits, dim=-1)
 
 
41
 
42
- # Defensive fallback: if anything still slips through as
43
- # nan/inf/negative, fall back to a safe uniform distribution
44
- # instead of crashing the whole generation call.
45
- probs = torch.nan_to_num(probs, nan=0.0, posinf=0.0, neginf=0.0)
46
- if probs.sum().item() <= 0:
47
- probs = torch.ones_like(probs) / probs.size(-1)
48
 
49
- next_token = torch.multinomial(probs, num_samples=1)
 
 
 
50
 
51
- input_ids = torch.cat([input_ids, next_token], dim=-1)
52
 
53
- if next_token.item() == tokenizer.eos_token_id:
54
- break
55
 
56
- return tokenizer.decode(input_ids[0], skip_special_tokens=True)
 
 
57
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
58
 
59
- with gr.Blocks(title="FastLLM (150M) Demo") as demo:
60
  gr.Markdown(
61
  """
62
- # FastLLM (150M) — Modern Causal Language Model
63
- A ~150M parameter decoder-only model (GQA, SwiGLU, RMSNorm, RoPE) from
64
- [devoppro/FastLLM](https://huggingface.co/devoppro/FastLLM), still training.
65
- Small model, so expect small-model quality — this is a demo, not a chatbot.
 
 
66
  """
 
 
 
67
  )
 
 
 
 
 
 
 
 
68
  with gr.Row():
69
- with gr.Column():
70
- prompt = gr.Textbox(
71
- label="Prompt",
72
- value="Once upon a time,",
73
- lines=4,
74
- )
75
- max_new_tokens = gr.Slider(16, 256, value=60, step=8, label="Max new tokens")
76
- temperature = gr.Slider(0.1, 1.5, value=0.7, step=0.05, label="Temperature")
77
- top_k = gr.Slider(0, 100, value=40, step=5, label="Top-k")
78
- run_btn = gr.Button("Generate", variant="primary")
79
- with gr.Column():
80
- output = gr.Textbox(label="Output", lines=12)
81
-
82
- run_btn.click(
83
- fn=generate,
84
- inputs=[prompt, max_new_tokens, temperature, top_k],
85
- outputs=output,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
86
  )
87
 
88
- gr.Examples(
89
- examples=[
90
- ["Once upon a time,", 60, 0.7, 40],
91
- ["Who are you?", 60, 0.7, 40],
92
- ["The most important thing about machine learning is", 60, 0.7, 40],
 
 
 
 
 
 
 
 
93
  ],
94
- inputs=[prompt, max_new_tokens, temperature, top_k],
95
  )
96
 
 
 
 
 
 
 
 
 
 
 
 
97
  if __name__ == "__main__":
98
- demo.queue().launch()
 
 
 
 
1
+ import os
2
+ import gc
3
  import torch
4
+ import gradio as gr
5
+
6
+ from transformers import AutoTokenizer, AutoModelForCausalLM
7
 
8
  MODEL_ID = "devoppro/FastLLM"
 
9
 
10
+ # ---------------------------------------------------------
11
+ # Environment
12
+ # ---------------------------------------------------------
13
+
14
+ os.environ.setdefault("HF_HOME", "/tmp/huggingface")
15
+ os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
16
+
17
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
18
+
19
+ print("=" * 60)
20
+ print("FastLLM")
21
+ print("=" * 60)
22
+ print(f"Model: {MODEL_ID}")
23
+ print(f"Device: {DEVICE}")
24
+
25
+ # ---------------------------------------------------------
26
+ # Load tokenizer
27
+ # ---------------------------------------------------------
28
+
29
+ print("Loading tokenizer...")
30
+
31
+ tokenizer = AutoTokenizer.from_pretrained(
32
  MODEL_ID,
 
33
  trust_remote_code=True,
34
  )
35
+
36
+ if tokenizer.pad_token is None:
37
+ tokenizer.pad_token = tokenizer.eos_token
38
+
39
+ print("Tokenizer loaded.")
40
+
41
+ # ---------------------------------------------------------
42
+ # Load model
43
+ # ---------------------------------------------------------
44
+
45
+ print("Loading FastLLM...")
46
+
47
+ model_kwargs = {
48
+ "trust_remote_code": True,
49
+ "low_cpu_mem_usage": True,
50
+ }
51
+
52
+ if DEVICE == "cuda":
53
+ model_kwargs["torch_dtype"] = torch.float16
54
+ else:
55
+ model_kwargs["torch_dtype"] = torch.float32
56
+
57
+ model = AutoModelForCausalLM.from_pretrained(
58
+ MODEL_ID,
59
+ **model_kwargs,
60
+ )
61
+
62
  model.eval()
63
+ model.to(DEVICE)
64
+
65
+ print("FastLLM loaded successfully.")
66
+ print("=" * 60)
67
+
68
+
69
+ # ---------------------------------------------------------
70
+ # Generation
71
+ # ---------------------------------------------------------
72
+
73
+ def generate_response(
74
+ message,
75
+ history,
76
+ temperature,
77
+ top_p,
78
+ max_new_tokens,
79
+ repetition_penalty,
80
+ ):
81
+ if not message or not message.strip():
82
+ return history, ""
83
+
84
+ try:
85
+ # Convert Gradio history to a simple conversation
86
+ prompt_parts = []
87
+
88
+ for item in history:
89
+ if isinstance(item, dict):
90
+ role = item.get("role")
91
+ content = item.get("content", "")
92
+
93
+ if role == "user":
94
+ prompt_parts.append(f"User: {content}")
95
+
96
+ elif role == "assistant":
97
+ prompt_parts.append(f"Assistant: {content}")
98
+
99
+ prompt_parts.append(f"User: {message}")
100
+ prompt_parts.append("Assistant:")
101
+
102
+ prompt = "\n".join(prompt_parts)
103
+
104
+ inputs = tokenizer(
105
+ prompt,
106
+ return_tensors="pt",
107
+ truncation=True,
108
+ max_length=2048,
109
+ )
110
+
111
+ input_ids = inputs["input_ids"].to(DEVICE)
112
+ attention_mask = inputs["attention_mask"].to(DEVICE)
113
+
114
+ with torch.inference_mode():
115
+ output = model.generate(
116
+ input_ids=input_ids,
117
+ attention_mask=attention_mask,
118
+ max_new_tokens=int(max_new_tokens),
119
+ temperature=float(temperature),
120
+ top_p=float(top_p),
121
+ repetition_penalty=float(repetition_penalty),
122
+ do_sample=True,
123
+ pad_token_id=tokenizer.pad_token_id,
124
+ eos_token_id=tokenizer.eos_token_id,
125
+ )
126
+
127
+ generated_tokens = output[0][input_ids.shape[-1]:]
128
+
129
+ response = tokenizer.decode(
130
+ generated_tokens,
131
+ skip_special_tokens=True,
132
+ ).strip()
133
+
134
+ if not response:
135
+ response = "FastLLM did not generate a response."
136
 
137
+ history = history + [
138
+ {"role": "user", "content": message},
139
+ {"role": "assistant", "content": response},
140
+ ]
141
 
142
+ # Release temporary tensors
143
+ del inputs
144
+ del input_ids
145
+ del attention_mask
146
+ del output
 
147
 
148
+ if DEVICE == "cuda":
149
+ torch.cuda.empty_cache()
 
150
 
151
+ gc.collect()
 
 
 
 
152
 
153
+ return history, ""
 
 
154
 
155
+ except Exception as e:
156
+ print("Generation error:")
157
+ print(repr(e))
158
 
159
+ error_message = f"❌ Generation error:\n\n`{str(e)}`"
 
 
 
 
 
160
 
161
+ history = history + [
162
+ {"role": "user", "content": message},
163
+ {"role": "assistant", "content": error_message},
164
+ ]
165
 
166
+ return history, ""
167
 
 
 
168
 
169
+ # ---------------------------------------------------------
170
+ # Clear conversation
171
+ # ---------------------------------------------------------
172
 
173
+ def clear_chat():
174
+ return []
175
+
176
+
177
+ # ---------------------------------------------------------
178
+ # UI
179
+ # ---------------------------------------------------------
180
+
181
+ css = """
182
+ .gradio-container {
183
+ max-width: 1100px !important;
184
+ margin: auto !important;
185
+ }
186
+
187
+ .title {
188
+ text-align: center;
189
+ margin-bottom: 4px;
190
+ }
191
+
192
+ .subtitle {
193
+ text-align: center;
194
+ opacity: 0.7;
195
+ margin-bottom: 20px;
196
+ }
197
+
198
+ .status {
199
+ text-align: center;
200
+ font-size: 13px;
201
+ opacity: 0.65;
202
+ }
203
+ """
204
+
205
+ with gr.Blocks(
206
+ css=css,
207
+ title="FastLLM",
208
+ ) as demo:
209
 
 
210
  gr.Markdown(
211
  """
212
+ # ⚡ FastLLM
213
+ """,
214
+ elem_classes=["title"],
215
+ )
216
+
217
+ gr.Markdown(
218
  """
219
+ A 150M parameter causal language model built from scratch by **devoppro**.
220
+ """,
221
+ elem_classes=["subtitle"],
222
  )
223
+
224
+ chatbot = gr.Chatbot(
225
+ label="FastLLM",
226
+ height=560,
227
+ type="messages",
228
+ bubble_full_width=False,
229
+ )
230
+
231
  with gr.Row():
232
+
233
+ message = gr.Textbox(
234
+ placeholder="Message FastLLM...",
235
+ label="",
236
+ scale=5,
237
+ lines=2,
238
+ )
239
+
240
+ send = gr.Button(
241
+ "Send",
242
+ variant="primary",
243
+ scale=1,
244
+ )
245
+
246
+ with gr.Row():
247
+
248
+ temperature = gr.Slider(
249
+ minimum=0.1,
250
+ maximum=2.0,
251
+ value=0.7,
252
+ step=0.05,
253
+ label="Temperature",
254
+ )
255
+
256
+ top_p = gr.Slider(
257
+ minimum=0.1,
258
+ maximum=1.0,
259
+ value=0.9,
260
+ step=0.05,
261
+ label="Top P",
262
+ )
263
+
264
+ max_tokens = gr.Slider(
265
+ minimum=16,
266
+ maximum=1024,
267
+ value=256,
268
+ step=16,
269
+ label="Max New Tokens",
270
+ )
271
+
272
+ repetition_penalty = gr.Slider(
273
+ minimum=1.0,
274
+ maximum=2.0,
275
+ value=1.05,
276
+ step=0.01,
277
+ label="Repetition Penalty",
278
+ )
279
+
280
+ with gr.Row():
281
+
282
+ clear = gr.Button(
283
+ "🗑️ Clear conversation"
284
+ )
285
+
286
+ gr.Markdown(
287
+ f"""
288
+ <div class="status">
289
+ Model: <b>{MODEL_ID}</b> · Device: <b>{DEVICE.upper()}</b>
290
+ </div>
291
+ """,
292
+ elem_classes=["status"],
293
+ )
294
+
295
+ # -----------------------------------------------------
296
+ # Events
297
+ # -----------------------------------------------------
298
+
299
+ send.click(
300
+ generate_response,
301
+ inputs=[
302
+ message,
303
+ chatbot,
304
+ temperature,
305
+ top_p,
306
+ max_tokens,
307
+ repetition_penalty,
308
+ ],
309
+ outputs=[
310
+ chatbot,
311
+ message,
312
+ ],
313
  )
314
 
315
+ message.submit(
316
+ generate_response,
317
+ inputs=[
318
+ message,
319
+ chatbot,
320
+ temperature,
321
+ top_p,
322
+ max_tokens,
323
+ repetition_penalty,
324
+ ],
325
+ outputs=[
326
+ chatbot,
327
+ message,
328
  ],
 
329
  )
330
 
331
+ clear.click(
332
+ clear_chat,
333
+ inputs=[],
334
+ outputs=[chatbot],
335
+ )
336
+
337
+
338
+ # ---------------------------------------------------------
339
+ # Launch
340
+ # ---------------------------------------------------------
341
+
342
  if __name__ == "__main__":
343
+ demo.launch(
344
+ server_name="0.0.0.0",
345
+ server_port=7860,
346
+ )