Arush kumar commited on
Commit
a098b7a
·
1 Parent(s): bab27e7

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +268 -0
app.py ADDED
@@ -0,0 +1,268 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import os
4
+ os.environ["KERAS_BACKEND"] = "jax"
5
+
6
+ import numpy as np
7
+ import jax
8
+ import keras
9
+ import gradio as gr
10
+ from pathlib import Path
11
+
12
+ from veylon_model import create_llm
13
+ from tokenizer import TokenizerWrapper
14
+ from config import (
15
+ CONTEXT,
16
+ vocab_size,
17
+ D_MODEL,
18
+ numberoflayers,
19
+ numberofheads,
20
+ d_Latent,
21
+ ffn_mult,
22
+ num_kv_heads,
23
+ swa_window,
24
+ )
25
+
26
+ # ============================================================
27
+ # Initialize (runs once)
28
+ # ============================================================
29
+
30
+ keras.mixed_precision.set_global_policy("mixed_bfloat16")
31
+
32
+ print(f"Backend: {keras.backend.backend()}")
33
+ print(f"JAX devices: {jax.devices()}")
34
+
35
+ # Load tokenizer
36
+ tokenizer = TokenizerWrapper("tokenizer.model")
37
+ assert tokenizer.vocab_size == vocab_size, (
38
+ f"Tokenizer vocab ({tokenizer.vocab_size}) != config vocab ({vocab_size})"
39
+ )
40
+ print(f"✓ Tokenizer loaded: {tokenizer.vocab_size} vocab")
41
+
42
+ # Build model
43
+ print("Building model...")
44
+ model = create_llm(
45
+ vocab_size=vocab_size,
46
+ d_model=D_MODEL,
47
+ n_layers=numberoflayers,
48
+ n_heads=numberofheads,
49
+ d_latent=d_Latent,
50
+ ffn_mult=ffn_mult,
51
+ max_seq_len=CONTEXT,
52
+ use_moe=False,
53
+ num_kv_heads=num_kv_heads,
54
+ swa_window=swa_window,
55
+ )
56
+
57
+ # Warmup
58
+ dummy = np.zeros((1, CONTEXT), dtype=np.int32)
59
+ _ = model(dummy, training=False)
60
+ print("✓ Model built successfully")
61
+
62
+ # Load weights
63
+ WEIGHTS_PATH = "veylon_final.weights.h5"
64
+ if Path(WEIGHTS_PATH).exists():
65
+ print(f"Loading weights from: {WEIGHTS_PATH}")
66
+ model.load_weights(WEIGHTS_PATH)
67
+ print("✓ Weights loaded successfully")
68
+ else:
69
+ print(f"WARNING: {WEIGHTS_PATH} not found. Using untrained model.")
70
+
71
+ print(f"✓ Model params: {model.count_params():,}\n")
72
+
73
+ # ============================================================
74
+ # Sampling
75
+ # ============================================================
76
+
77
+ def sample_from_logits(
78
+ logits: np.ndarray,
79
+ temperature: float = 0.8,
80
+ top_k: int = 50,
81
+ ) -> int:
82
+ """NumPy-only sampling."""
83
+ logits = np.array(logits, dtype=np.float32, copy=True)
84
+
85
+ if temperature > 0:
86
+ logits = logits / float(max(temperature, 1e-8))
87
+
88
+ if top_k > 0:
89
+ k = min(int(top_k), logits.shape[-1])
90
+ row = logits[0]
91
+ top_indices = np.argpartition(row, -k)[-k:]
92
+ filtered = np.full_like(row, -np.inf)
93
+ filtered[top_indices] = row[top_indices]
94
+ logits[0] = filtered
95
+
96
+ row = logits[0]
97
+ row = row - np.max(row)
98
+ probs = np.exp(row)
99
+ probs = probs / probs.sum()
100
+
101
+ return int(np.random.choice(len(probs), p=probs))
102
+
103
+ # ============================================================
104
+ # Generation function
105
+ # ============================================================
106
+
107
+ def generate(
108
+ prompt: str,
109
+ max_new_tokens: int = 64,
110
+ temperature: float = 0.8,
111
+ top_k: int = 50,
112
+ ) -> str:
113
+ """
114
+ Generate text from a prompt using Veylon.
115
+
116
+ Args:
117
+ prompt: Input text
118
+ max_new_tokens: Maximum tokens to generate
119
+ temperature: Sampling temperature (0.1-2.0)
120
+ top_k: Top-K sampling cutoff
121
+
122
+ Returns:
123
+ Generated text
124
+ """
125
+ try:
126
+ # Encode prompt
127
+ tokens = tokenizer.encode(
128
+ prompt,
129
+ add_bos=True,
130
+ add_eos=False,
131
+ )
132
+
133
+ if len(tokens) == 0:
134
+ tokens = [tokenizer.bos_id if hasattr(tokenizer, "bos_id") else 1]
135
+
136
+ tokens = tokens[-CONTEXT:]
137
+
138
+ # Prefill phase
139
+ prompt_ids = np.array([tokens], dtype=np.int32)
140
+ logits, cache_k, cache_v = model.generate_step(
141
+ prompt_ids,
142
+ cache_k=None,
143
+ cache_v=None,
144
+ cache_pos=0,
145
+ )
146
+
147
+ next_token = sample_from_logits(
148
+ np.array(logits[:, -1, :], dtype=np.float32, copy=True),
149
+ temperature=temperature,
150
+ top_k=top_k,
151
+ )
152
+ tokens.append(next_token)
153
+
154
+ # Decoding phase (token-by-token)
155
+ if next_token != tokenizer.eos_id and len(tokens) < CONTEXT:
156
+ cache_pos = len(prompt_ids[0])
157
+
158
+ for _ in range(max_new_tokens - 1):
159
+ next_input = np.array([[next_token]], dtype=np.int32)
160
+
161
+ logits, cache_k, cache_v = model.generate_step(
162
+ next_input,
163
+ cache_k=cache_k,
164
+ cache_v=cache_v,
165
+ cache_pos=cache_pos,
166
+ )
167
+
168
+ cache_pos += 1
169
+
170
+ next_token = sample_from_logits(
171
+ np.array(logits[:, -1, :], dtype=np.float32, copy=True),
172
+ temperature=temperature,
173
+ top_k=top_k,
174
+ )
175
+ tokens.append(next_token)
176
+
177
+ if next_token == tokenizer.eos_id:
178
+ break
179
+
180
+ if len(tokens) >= CONTEXT:
181
+ break
182
+
183
+ # Decode output
184
+ generated_text = tokenizer.decode(tokens)
185
+ return generated_text
186
+
187
+ except Exception as e:
188
+ return f"Error: {str(e)}"
189
+
190
+ # ============================================================
191
+ # Gradio UI
192
+ # ============================================================
193
+
194
+ def main():
195
+ with gr.Blocks(title="Veylon Alpha") as demo:
196
+ gr.Markdown("""
197
+ # 🚀 Veylon Alpha - 10M LLM
198
+
199
+ A small transformer model trained on clean data.
200
+ Enter a prompt and watch it generate text.
201
+ """)
202
+
203
+ with gr.Row():
204
+ with gr.Column(scale=2):
205
+ prompt = gr.Textbox(
206
+ label="Prompt",
207
+ placeholder="Once upon a time",
208
+ lines=3,
209
+ value="Once upon a time"
210
+ )
211
+
212
+ with gr.Row():
213
+ max_tokens = gr.Slider(
214
+ label="Max tokens",
215
+ minimum=10,
216
+ maximum=256,
217
+ value=64,
218
+ step=10,
219
+ )
220
+ temperature = gr.Slider(
221
+ label="Temperature",
222
+ minimum=0.1,
223
+ maximum=2.0,
224
+ value=0.8,
225
+ step=0.1,
226
+ )
227
+ top_k = gr.Slider(
228
+ label="Top-K",
229
+ minimum=1,
230
+ maximum=100,
231
+ value=50,
232
+ step=1,
233
+ )
234
+
235
+ generate_btn = gr.Button("Generate", variant="primary", size="lg")
236
+
237
+ with gr.Column(scale=1):
238
+ info = gr.Markdown(f"""
239
+ **Model Info**
240
+
241
+ - Parameters: {model.count_params():,}
242
+ - Context: {CONTEXT} tokens
243
+ - Vocab: {vocab_size}
244
+ - Architecture: Transformer + GQA
245
+
246
+ **Tips**
247
+ - Higher temp = more creative
248
+ - Lower temp = more deterministic
249
+ - Top-K = diversity control
250
+ """)
251
+
252
+ output = gr.Textbox(
253
+ label="Generated Output",
254
+ lines=8,
255
+ interactive=False
256
+ )
257
+
258
+ # Connect
259
+ generate_btn.click(
260
+ fn=generate,
261
+ inputs=[prompt, max_tokens, temperature, top_k],
262
+ outputs=output
263
+ )
264
+
265
+ demo.launch(share=False, server_name="0.0.0.0", server_port=7860)
266
+
267
+ if __name__ == "__main__":
268
+ main()