Arush kumar commited on
Commit
b011ee8
·
1 Parent(s): cd97d62

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +253 -503
app.py CHANGED
@@ -1,521 +1,271 @@
1
  from __future__ import annotations
2
 
 
 
 
3
  import numpy as np
4
- import keras
5
- from keras import layers, ops
6
  import jax
7
-
8
- from veylon_attention import flash_splash_attention, decode_swa
9
-
10
- try:
11
- from config import (
12
- CONTEXT,
13
- vocab_size as Vocab_size,
14
- D_MODEL,
15
- numberoflayers,
16
- numberofheads,
17
- d_Latent,
18
- ffn_mult,
19
- swa_window,
20
- num_kv_heads,
21
- use_moe,
22
- moe_num_experts,
23
- moe_top_k,
24
- USE_REMAT,
25
- MAX_GEN_TOKENS,
26
- )
27
- except Exception:
28
- CONTEXT = 2048
29
- Vocab_size = 32000
30
- D_MODEL = 512
31
- numberoflayers = 8
32
- numberofheads = 8
33
- d_Latent = 128
34
- ffn_mult = 3.5
35
- swa_window = 1024
36
- num_kv_heads = 2
37
- use_moe = False
38
- moe_num_experts = 8
39
- moe_top_k = 2
40
- USE_REMAT = True
41
- MAX_GEN_TOKENS = 256
42
-
43
- # NOTE: mixed precision policy is set by train.py / finetune.py before model
44
- # creation. Do NOT set it here — importing this module must not have side
45
- # effects.
46
-
47
-
48
- # ─────────────────────────────────────────────────────────────────────────────
49
- # Layers
50
- # ─────────────────────────────────────────────────────────────────────────────
51
-
52
- @keras.saving.register_keras_serializable()
53
- class RMSNorm(layers.Layer):
54
- def __init__(self, epsilon=1e-5, **kwargs):
55
- super().__init__(**kwargs)
56
- self.epsilon = epsilon
57
-
58
- def build(self, input_shape):
59
- self.weight = self.add_weight(shape=(input_shape[-1],), initializer='ones', name='gamma')
60
-
61
- def call(self, x):
62
- x_fp32 = ops.cast(x, 'float32')
63
- rms = ops.sqrt(ops.mean(ops.square(x_fp32), axis=-1, keepdims=True) + self.epsilon)
64
- out = x_fp32 / rms
65
- return ops.cast(out, x.dtype) * self.weight
66
-
67
- def get_config(self):
68
- cfg = super().get_config()
69
- cfg.update({'epsilon': self.epsilon})
70
- return cfg
71
-
72
-
73
- @keras.saving.register_keras_serializable()
74
- class RotaryEmbedding(layers.Layer):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75
  """
76
- RoPE with `max_seq_len + gen_headroom` so a full-context prompt can still
77
- decode up to `gen_headroom` tokens without overflowing the cos/sin table.
78
-
79
- `max_seq_len` is preserved in the config so checkpoints remain compatible —
80
- only the (trainable=False) table shapes grow.
 
 
 
 
 
81
  """
 
 
 
 
 
 
 
 
 
 
 
 
82
 
83
- def __init__(self, max_seq_len, dim, theta=10000.0, gen_headroom=0, **kwargs):
84
- super().__init__(**kwargs)
85
- if dim % 2 != 0:
86
- raise ValueError('RotaryEmbedding dim must be even.')
87
- self.max_seq_len = max_seq_len
88
- self.dim = dim
89
- self.theta = theta
90
- self.gen_headroom = int(gen_headroom)
91
- self.table_size = max_seq_len + self.gen_headroom
92
-
93
- def build(self, input_shape):
94
- half = self.dim // 2
95
- inv_freq = 1.0 / (self.theta ** (np.arange(0, self.dim, 2).astype(np.float32) / self.dim))
96
- positions = np.arange(self.table_size, dtype=np.float32)
97
- freqs = positions[:, None] * inv_freq[None, :]
98
- self.cos = self.add_weight(
99
- shape=(self.table_size, half),
100
- initializer=keras.initializers.Constant(np.cos(freqs)),
101
- trainable=False, dtype='float32', name='cos_table',
102
  )
103
- self.sin = self.add_weight(
104
- shape=(self.table_size, half),
105
- initializer=keras.initializers.Constant(np.sin(freqs)),
106
- trainable=False, dtype='float32', name='sin_table',
 
107
  )
 
 
 
 
 
108
 
109
- def call(self, x, offset=0):
110
- seq_len = x.shape[1]
111
- if seq_len is None:
112
- seq_len = ops.shape(x)[1]
113
- half = self.dim // 2
114
-
115
- # JIT FIX: `offset` may be a traced jnp array (e.g. cache_pos passed
116
- # in under jax.jit for compiled autoregressive decoding), not a
117
- # Python int. The original code did `if offset + seq_len >
118
- # self.table_size: raise ...` and `self.cos[offset:offset+seq_len]`
119
- # — both are illegal under tracing: you cannot branch a Python `if`
120
- # on a traced value, and Python slice syntax requires a static start
121
- # index. We replace the bounds check with a debug-only assertion
122
- # that only fires when offset is a concrete Python int/np scalar
123
- # (i.e. eager calls, not jitted ones — jitted calls trust the caller
124
- # to pass a valid offset, same as the KV-cache contract elsewhere in
125
- # this file), and replace the slice with ops.slice /
126
- # dynamic_slice_in_dim, which correctly handles a traced start index.
127
- if isinstance(offset, int):
128
- if offset + seq_len > self.table_size:
129
- raise ValueError(
130
- f'RoPE table too small: offset={offset}, seq_len={seq_len}, '
131
- f'max={self.table_size} (CONTEXT={self.max_seq_len} + gen_headroom={self.gen_headroom})'
132
  )
133
 
134
- offset = ops.cast(offset, 'int32')
135
- cos = ops.slice(self.cos, [offset, 0], [seq_len, half])
136
- sin = ops.slice(self.sin, [offset, 0], [seq_len, half])
137
- cos = ops.cast(ops.reshape(cos, (1, seq_len, 1, half)), x.dtype)
138
- sin = ops.cast(ops.reshape(sin, (1, seq_len, 1, half)), x.dtype)
139
- x1 = x[..., :half]
140
- x2 = x[..., half:]
141
- return ops.concatenate([x1 * cos - x2 * sin, x1 * sin + x2 * cos], axis=-1)
142
-
143
- def get_config(self):
144
- cfg = super().get_config()
145
- cfg.update({
146
- 'max_seq_len': self.max_seq_len,
147
- 'dim': self.dim,
148
- 'theta': self.theta,
149
- 'gen_headroom': self.gen_headroom,
150
- })
151
- return cfg
152
-
153
-
154
- @keras.saving.register_keras_serializable()
155
- class SwiGLUFFN(layers.Layer):
156
- def __init__(self, d_model, hidden_mult=3.5, **kwargs):
157
- super().__init__(**kwargs)
158
- self.d_model_arg = d_model
159
- self.hidden_mult = hidden_mult
160
- self.hidden_dim = int(d_model * hidden_mult * 2 / 3)
161
- self.hidden_dim = ((self.hidden_dim + 63) // 64) * 64
162
-
163
- def build(self, input_shape):
164
- d_model = input_shape[-1]
165
- self.gate_up_proj = self.add_weight(shape=(d_model, 2 * self.hidden_dim), initializer='glorot_uniform', name='gate_up_proj')
166
- self.down_proj = self.add_weight(shape=(self.hidden_dim, d_model), initializer='glorot_uniform', name='down_proj')
167
-
168
- def call(self, x, training=False):
169
- gate_up = ops.matmul(x, self.gate_up_proj)
170
- gate, up = ops.split(gate_up, 2, axis=-1)
171
- return ops.matmul(ops.silu(gate) * up, self.down_proj)
172
-
173
- def get_config(self):
174
- cfg = super().get_config()
175
- cfg.update({'d_model': self.d_model_arg, 'hidden_mult': self.hidden_mult})
176
- return cfg
177
-
178
-
179
- @keras.saving.register_keras_serializable()
180
- class MoE_FFN(layers.Layer):
181
- """
182
- Mixture-of-Experts FFN (top-k routing, masked dispatch).
183
 
184
- For each routing slot k in [0, top_k) we run each expert once over the
185
- full batch, zeroing inputs the expert was not selected for, then accumulate
186
- the gated expert output. The load-balancing loss uses top-1 argmax routing
187
- for the `f` term (standard Switch-Transformer formulation); the previous
188
- code used top-k membership which over-counted and produced the wrong loss.
189
- """
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
190
 
191
- def __init__(self, d_model, num_experts=8, top_k=2, hidden_mult=3.5, **kwargs):
192
- super().__init__(**kwargs)
193
- self.d_model_arg = d_model
194
- self.num_experts = num_experts
195
- self.top_k = top_k
196
- self.hidden_mult = hidden_mult
197
- self.experts = [SwiGLUFFN(d_model, hidden_mult) for _ in range(num_experts)]
198
- self.router = layers.Dense(num_experts, use_bias=False)
199
-
200
- def _load_balancing_loss(self, router_logits, top_idx):
201
- # Standard Switch-Transformer aux loss. `f` uses the argmax routing
202
- # decision (top_idx[..., 0]); `p` uses the router softmax mean.
203
- router_probs = ops.softmax(router_logits, axis=-1)
204
- mask = ops.one_hot(top_idx[..., 0], self.num_experts) # [..., num_experts]
205
- f = ops.mean(ops.cast(mask, 'float32'), axis=tuple(range(mask.ndim - 1)))
206
- p = ops.mean(router_probs, axis=tuple(range(router_probs.ndim - 1)))
207
- return self.num_experts * ops.sum(f * p)
208
-
209
- def call(self, x, training=False):
210
- # x: [B, S, D]
211
- B = ops.shape(x)[0]
212
- S = ops.shape(x)[1]
213
- D = x.shape[-1]
214
-
215
- router_logits = self.router(x) # [B, S, E]
216
- top_logits, top_idx = ops.top_k(router_logits, self.top_k) # [B, S, k]
217
- top_weights = ops.softmax(top_logits, axis=-1) # [B, S, k]
218
-
219
- if training:
220
- self.add_loss(self._load_balancing_loss(router_logits, top_idx))
221
-
222
- out = ops.zeros_like(x)
223
- for kk in range(self.top_k):
224
- idx_k = top_idx[..., kk] # [B, S]
225
- w_k = top_weights[..., kk][..., None] # [B, S, 1]
226
- for e in range(self.num_experts):
227
- mask_e = ops.cast(idx_k == e, x.dtype)[..., None] # [B, S, 1]
228
- # mask_e is always a tensor; skip the useless ops.is_tensor guard
229
- # (it never triggered — removed FIX BUG11)
230
- # Only run the expert where it is actually selected.
231
- x_e = x * mask_e
232
- expert_out = self.experts[e](x_e, training=training)
233
- out = out + expert_out * w_k * mask_e
234
- return out
235
-
236
- def get_config(self):
237
- cfg = super().get_config()
238
- cfg.update({
239
- 'd_model': self.d_model_arg, 'num_experts': self.num_experts,
240
- 'top_k': self.top_k, 'hidden_mult': self.hidden_mult,
241
- })
242
- return cfg
243
-
244
-
245
- @keras.saving.register_keras_serializable()
246
- class MLAttention(layers.Layer):
247
- def __init__(self, d_model, n_heads, d_latent, max_seq_len, num_kv_heads=2,
248
- swa_window=1024, attn_dropout=0.0, gen_headroom=0, **kwargs):
249
- super().__init__(**kwargs)
250
- if d_model % n_heads != 0:
251
- raise ValueError('d_model must be divisible by n_heads')
252
- if n_heads % num_kv_heads != 0:
253
- raise ValueError('n_heads must be divisible by num_kv_heads')
254
- self.d_model = d_model
255
- self.n_heads = n_heads
256
- self.num_kv_heads = num_kv_heads
257
- self.group_size = n_heads // num_kv_heads
258
- self.d_head = d_model // n_heads
259
- self.d_latent = d_latent
260
- self.max_seq_len = max_seq_len
261
- self.swa_window = swa_window
262
- self.gen_headroom = int(gen_headroom)
263
- self.dropout = layers.Dropout(attn_dropout)
264
- self.rope = RotaryEmbedding(max_seq_len, d_model // n_heads, gen_headroom=int(gen_headroom)) # FIX BUG10: must be in __init__ for Keras tracking
265
-
266
- def build(self, input_shape):
267
- self.W_qc = self.add_weight(shape=(self.d_model, self.d_model + self.d_latent), initializer='glorot_uniform', name='W_qc')
268
- self.W_kv = self.add_weight(shape=(self.d_latent, self.num_kv_heads * 2 * self.d_head), initializer='glorot_uniform', name='W_kv')
269
- self.W_o = self.add_weight(shape=(self.d_model, self.d_model), initializer='glorot_uniform', name='Wo')
270
-
271
- def _project_kv(self, c):
272
- kv = ops.matmul(c, self.W_kv)
273
- return ops.split(kv, 2, axis=-1)
274
-
275
- def call(self, x, training=False):
276
- B = ops.shape(x)[0]
277
- S = ops.shape(x)[1]
278
- qc = ops.matmul(x, self.W_qc)
279
- q_proj, c = ops.split(qc, [self.d_model], axis=-1)
280
- q = ops.reshape(q_proj, (B, S, self.n_heads, self.d_head))
281
- q = self.rope(q, offset=0)
282
- q = ops.transpose(q, (0, 2, 1, 3))
283
- k, v = self._project_kv(c)
284
- k = ops.reshape(k, (B, S, self.num_kv_heads, self.d_head))
285
- v = ops.reshape(v, (B, S, self.num_kv_heads, self.d_head))
286
- k = self.rope(k, offset=0)
287
- k = ops.transpose(k, (0, 2, 1, 3))
288
- v = ops.transpose(v, (0, 2, 1, 3))
289
-
290
- # Attention remat is REQUIRED during training: it is the per-block
291
- # jax.checkpoint(process_block) that stops XLA from staging the hidden
292
- # [S, B, Hkv, W, D] buffer (≈8.6 GB here) inside the fori_loop. Block-
293
- # level remat does NOT prevent that forward-time allocation, so this
294
- # must follow `training` regardless of config.USE_REMAT.
295
- out = flash_splash_attention(
296
- q, k, v,
297
- window_size=min(self.swa_window, self.max_seq_len),
298
- backend=jax.default_backend(),
299
- use_gqa=True,
300
- use_remat=training,
301
- )
302
- out = ops.transpose(out, (0, 2, 1, 3))
303
- out = ops.reshape(out, (B, S, self.d_model))
304
- out = self.dropout(out, training=training)
305
- return ops.matmul(out, self.W_o)
306
-
307
- def generate_step(self, x, cache_k=None, cache_v=None, cache_pos=0):
308
- B = ops.shape(x)[0]
309
- S = ops.shape(x)[1]
310
-
311
- qc = ops.matmul(x, self.W_qc)
312
- q_proj, c = ops.split(qc, [self.d_model], axis=-1)
313
-
314
- q = ops.reshape(q_proj, (B, S, self.n_heads, self.d_head))
315
- q = self.rope(q, offset=cache_pos)
316
- q = ops.transpose(q, (0, 2, 1, 3))
317
-
318
- k, v = self._project_kv(c)
319
- k = ops.reshape(k, (B, S, self.num_kv_heads, self.d_head))
320
- v = ops.reshape(v, (B, S, self.num_kv_heads, self.d_head))
321
- k = self.rope(k, offset=cache_pos)
322
- k = ops.transpose(k, (0, 2, 1, 3))
323
- v = ops.transpose(v, (0, 2, 1, 3))
324
-
325
- if cache_k is None:
326
- # Prefill: no remat (inference only), no window trim needed.
327
- out = flash_splash_attention(
328
- q, k, v,
329
- window_size=min(self.swa_window, self.max_seq_len),
330
- backend=jax.default_backend(),
331
- use_gqa=True,
332
- use_remat=False,
333
- )
334
- new_k = k[:, :, -self.swa_window:, :]
335
- new_v = v[:, :, -self.swa_window:, :]
336
- else:
337
- if S != 1:
338
- raise ValueError(f'generate_step with cache expects S=1, got S={S}')
339
-
340
- k = ops.concatenate([cache_k, k], axis=2)
341
- v = ops.concatenate([cache_v, v], axis=2)
342
-
343
- k = k[:, :, -self.swa_window:, :]
344
- v = v[:, :, -self.swa_window:, :]
345
-
346
- out = decode_swa(q, k, v)
347
- new_k = k
348
- new_v = v
349
-
350
- out = ops.transpose(out, (0, 2, 1, 3))
351
- out = ops.reshape(out, (B, S, self.d_model))
352
- out = ops.matmul(out, self.W_o)
353
- return out, new_k, new_v
354
-
355
- def get_config(self):
356
- cfg = super().get_config()
357
- cfg.update({
358
- 'd_model': self.d_model, 'n_heads': self.n_heads,
359
- 'num_kv_heads': self.num_kv_heads, 'd_latent': self.d_latent,
360
- 'max_seq_len': self.max_seq_len, 'swa_window': self.swa_window,
361
- 'attn_dropout': self.dropout.rate, 'gen_headroom': self.gen_headroom,
362
- })
363
- return cfg
364
-
365
-
366
- @keras.saving.register_keras_serializable()
367
- class TransformerBlock(layers.Layer):
368
- def __init__(self, d_model, n_heads, d_latent, ffn_layer, max_seq_len,
369
- num_kv_heads=2, swa_window=1024, use_remat=True, gen_headroom=0, **kwargs):
370
- super().__init__(**kwargs)
371
- self.d_model = d_model
372
- self.n_heads = n_heads
373
- self.d_latent = d_latent
374
- self.max_seq_len = max_seq_len
375
- self.num_kv_heads = num_kv_heads
376
- self.swa_window = swa_window
377
- self.use_remat = use_remat
378
- self.gen_headroom = int(gen_headroom)
379
- self.ffn = keras.saving.deserialize_keras_object(ffn_layer) if isinstance(ffn_layer, dict) else ffn_layer
380
- self.norm1 = RMSNorm()
381
- self.norm2 = RMSNorm()
382
- self.attn = MLAttention(
383
- d_model, n_heads, d_latent, max_seq_len,
384
- num_kv_heads=num_kv_heads, swa_window=swa_window, gen_headroom=self.gen_headroom,
385
  )
386
 
387
- def call(self, x, training=False):
388
- # Block-level activation/gradient checkpointing (optional, via
389
- # config.USE_REMAT). NOTE: this is independent of attention's own
390
- # remat, which MUST stay on whenever training=True — that's what stops
391
- # XLA staging the hidden [S,B,Hkv,W,D] buffer (the OOM we hit).
392
- if training and self.use_remat:
393
- def _fwd(x_in):
394
- attn_out = self.attn(self.norm1(x_in), training=True)
395
- ffn_out = self.ffn(self.norm2(x_in), training=True)
396
- return attn_out + ffn_out
397
- return x + jax.checkpoint(_fwd)(x)
398
- a = self.attn(self.norm1(x), training=training)
399
- f = self.ffn(self.norm2(x), training=training)
400
- return x + a + f
401
-
402
- def generate_step(self, x, cache_k=None, cache_v=None, cache_pos=0):
403
- attn_out, nck, ncv = self.attn.generate_step(
404
- self.norm1(x), cache_k=cache_k, cache_v=cache_v, cache_pos=cache_pos,
405
  )
406
- x = x + attn_out
407
- x = x + self.ffn(self.norm2(x), training=False)
408
- return x, nck, ncv
409
-
410
- def get_config(self):
411
- cfg = super().get_config()
412
- cfg.update({
413
- 'd_model': self.d_model, 'n_heads': self.n_heads, 'd_latent': self.d_latent,
414
- 'ffn_layer': keras.saving.serialize_keras_object(self.ffn),
415
- 'max_seq_len': self.max_seq_len, 'num_kv_heads': self.num_kv_heads,
416
- 'swa_window': self.swa_window, 'use_remat': self.use_remat,
417
- 'gen_headroom': self.gen_headroom,
418
- })
419
- return cfg
420
-
421
-
422
- @keras.saving.register_keras_serializable()
423
- class VeylonModel(keras.Model):
424
- def __init__(self, vocab_size, d_model, n_layers, n_heads, d_latent, ffn_mult,
425
- max_seq_len, use_moe=False, moe_num_experts=8, moe_top_k=2,
426
- num_kv_heads=2, swa_window=1024, use_remat=True, gen_headroom=0, **kwargs):
427
- super().__init__(**kwargs)
428
- self.vocab_size = vocab_size
429
- self.d_model = d_model
430
- self.n_layers = n_layers
431
- self.n_heads = n_heads
432
- self.d_latent = d_latent
433
- self.ffn_mult = ffn_mult
434
- self.max_seq_len = max_seq_len
435
- self.use_moe = use_moe
436
- self.moe_num_experts = moe_num_experts
437
- self.moe_top_k = moe_top_k
438
- self.num_kv_heads = num_kv_heads
439
- self.swa_window = swa_window
440
- self.use_remat = use_remat
441
- self.gen_headroom = int(gen_headroom)
442
- self.embedding = layers.Embedding(vocab_size, d_model, name='token_embedding')
443
- self.blocks = []
444
- for i in range(n_layers):
445
- ffn = (MoE_FFN(d_model, moe_num_experts, moe_top_k, ffn_mult)
446
- if use_moe else SwiGLUFFN(d_model, ffn_mult))
447
- self.blocks.append(TransformerBlock(
448
- d_model, n_heads, d_latent, ffn, max_seq_len,
449
- num_kv_heads=num_kv_heads, swa_window=swa_window,
450
- use_remat=use_remat, gen_headroom=self.gen_headroom, name=f'block_{i}',
451
- ))
452
- self.norm = RMSNorm()
453
-
454
- def call(self, inputs, training=False):
455
- x = self.embedding(inputs)
456
- for block in self.blocks:
457
- x = block(x, training=training)
458
- x = self.norm(x)
459
- embedding_weights = self.embedding.embeddings
460
- logits = ops.matmul(x, ops.transpose(embedding_weights))
461
- return ops.cast(logits, 'float32')
462
-
463
- def generate_step(self, inputs, cache_k=None, cache_v=None, cache_pos=0):
464
- x = self.embedding(inputs)
465
- new_cache_k = []
466
- new_cache_v = []
467
-
468
- if cache_k is None:
469
- cache_k = [None] * len(self.blocks)
470
- cache_v = [None] * len(self.blocks)
471
-
472
- for i, block in enumerate(self.blocks):
473
- x, nck, ncv = block.generate_step(
474
- x, cache_k=cache_k[i], cache_v=cache_v[i], cache_pos=cache_pos,
475
- )
476
- new_cache_k.append(nck)
477
- new_cache_v.append(ncv)
478
-
479
- x = self.norm(x)
480
- embedding_weights = self.embedding.embeddings
481
- logits = ops.matmul(x, ops.transpose(embedding_weights))
482
- logits = ops.cast(logits, 'float32')
483
- return logits, new_cache_k, new_cache_v
484
-
485
- def get_config(self):
486
- cfg = super().get_config()
487
- cfg.update({
488
- 'vocab_size': self.vocab_size, 'd_model': self.d_model,
489
- 'n_layers': self.n_layers, 'n_heads': self.n_heads,
490
- 'd_latent': self.d_latent, 'ffn_mult': self.ffn_mult,
491
- 'max_seq_len': self.max_seq_len, 'use_moe': self.use_moe,
492
- 'moe_num_experts': self.moe_num_experts, 'moe_top_k': self.moe_top_k,
493
- 'num_kv_heads': self.num_kv_heads, 'swa_window': self.swa_window,
494
- 'use_remat': self.use_remat, 'gen_headroom': self.gen_headroom,
495
- })
496
- return cfg
497
-
498
-
499
- # ─────────────────────────────────────────────────────────────────────────────
500
- # Factory
501
- # ─────────────────────────────────────────────────────────────────────────────
502
-
503
- def create_llm(
504
- vocab_size=Vocab_size, d_model=D_MODEL, n_layers=numberoflayers, n_heads=numberofheads,
505
- d_latent=d_Latent, ffn_mult=ffn_mult, max_seq_len=CONTEXT, use_moe=use_moe,
506
- moe_num_experts=moe_num_experts, moe_top_k=moe_top_k, num_kv_heads=num_kv_heads,
507
- swa_window=swa_window, use_remat=USE_REMAT, gen_headroom=MAX_GEN_TOKENS,
508
- ):
509
- """
510
- Build a VeylonModel.
511
 
512
- `gen_headroom` sizes the RoPE table to (max_seq_len + gen_headroom) so that
513
- a full-length prompt can still generate `gen_headroom` tokens during
514
- inference without "RoPE table too small" errors. Defaults to MAX_GEN_TOKENS.
515
- """
516
- return VeylonModel(
517
- vocab_size=vocab_size, d_model=d_model, n_layers=n_layers, n_heads=n_heads,
518
- d_latent=d_latent, ffn_mult=ffn_mult, max_seq_len=max_seq_len, use_moe=use_moe,
519
- moe_num_experts=moe_num_experts, moe_top_k=moe_top_k, num_kv_heads=num_kv_heads,
520
- swa_window=swa_window, use_remat=use_remat, gen_headroom=gen_headroom,
521
- )
 
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 Preview - 10M LLM
198
+ # Made by Arush Kumar
199
+ A student dev!
200
+ A small transformer model trained on clean data.
201
+ Enter a prompt and watch it generate text.
202
+
203
+ """)
204
+
205
+ with gr.Row():
206
+ with gr.Column(scale=2):
207
+ prompt = gr.Textbox(
208
+ label="Prompt",
209
+ placeholder="Once upon a time",
210
+ lines=3,
211
+ value="Once upon a time"
212
+ )
213
 
214
+ with gr.Row():
215
+ max_tokens = gr.Slider(
216
+ label="Max tokens",
217
+ minimum=10,
218
+ maximum=256,
219
+ value=64,
220
+ step=10,
221
+ )
222
+ temperature = gr.Slider(
223
+ label="Temperature",
224
+ minimum=0.1,
225
+ maximum=2.0,
226
+ value=0.8,
227
+ step=0.1,
228
+ )
229
+ top_k = gr.Slider(
230
+ label="Top-K",
231
+ minimum=1,
232
+ maximum=100,
233
+ value=50,
234
+ step=1,
235
+ )
236
+
237
+ generate_btn = gr.Button("Generate", variant="primary", size="lg")
238
+
239
+ with gr.Column(scale=1):
240
+ info = gr.Markdown(f"""
241
+ **Model Info**
242
+
243
+ - Parameters: {model.count_params():,}
244
+ - Context: {CONTEXT} tokens
245
+ - Vocab: {vocab_size}
246
+ - Architecture: Transformer + GQA
247
+
248
+ **Tips**
249
+ - Higher temp = more creative
250
+ - Lower temp = more deterministic
251
+ - Top-K = diversity control
252
+ """)
253
+
254
+ output = gr.Textbox(
255
+ label="Generated Output",
256
+ lines=8,
257
+ interactive=False
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
258
  )
259
 
260
+ # Connect
261
+ generate_btn.click(
262
+ fn=generate,
263
+ inputs=[prompt, max_tokens, temperature, top_k],
264
+ outputs=output,
265
+ api_name="generate"
 
 
 
 
 
 
 
 
 
 
 
 
266
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
267
 
268
+ demo.launch(share=False, server_name="0.0.0.0", server_port=7860)
269
+
270
+ if __name__ == "__main__":
271
+ main()