Pragya / app.py
Arush kumar
Update app.py
b011ee8
Raw History Blame
7.59 kB
from __future__ import annotations
import os
os.environ["KERAS_BACKEND"] = "jax"
import numpy as np
import jax
import keras
import gradio as gr
from pathlib import Path
from veylon_model import create_llm
from tokenizer import TokenizerWrapper
from config import (
CONTEXT,
vocab_size,
D_MODEL,
numberoflayers,
numberofheads,
d_Latent,
ffn_mult,
num_kv_heads,
swa_window,
)
# ============================================================
# Initialize (runs once)
# ============================================================
keras.mixed_precision.set_global_policy("mixed_bfloat16")
print(f"Backend: {keras.backend.backend()}")
print(f"JAX devices: {jax.devices()}")
# Load tokenizer
tokenizer = TokenizerWrapper("tokenizer.model")
assert tokenizer.vocab_size == vocab_size, (
f"Tokenizer vocab ({tokenizer.vocab_size}) != config vocab ({vocab_size})"
)
print(f"✓ Tokenizer loaded: {tokenizer.vocab_size} vocab")
# Build model
print("Building model...")
model = create_llm(
vocab_size=vocab_size,
d_model=D_MODEL,
n_layers=numberoflayers,
n_heads=numberofheads,
d_latent=d_Latent,
ffn_mult=ffn_mult,
max_seq_len=CONTEXT,
use_moe=False,
num_kv_heads=num_kv_heads,
swa_window=swa_window,
)
# Warmup
dummy = np.zeros((1, CONTEXT), dtype=np.int32)
_ = model(dummy, training=False)
print("✓ Model built successfully")
# Load weights
WEIGHTS_PATH = "veylon_final.weights.h5"
if Path(WEIGHTS_PATH).exists():
print(f"Loading weights from: {WEIGHTS_PATH}")
model.load_weights(WEIGHTS_PATH)
print("✓ Weights loaded successfully")
else:
print(f"WARNING: {WEIGHTS_PATH} not found. Using untrained model.")
print(f"✓ Model params: {model.count_params():,}\n")
# ============================================================
# Sampling
# ============================================================
def sample_from_logits(
logits: np.ndarray,
temperature: float = 0.8,
top_k: int = 50,
) -> int:
"""NumPy-only sampling."""
logits = np.array(logits, dtype=np.float32, copy=True)
if temperature > 0:
logits = logits / float(max(temperature, 1e-8))
if top_k > 0:
k = min(int(top_k), logits.shape[-1])
row = logits[0]
top_indices = np.argpartition(row, -k)[-k:]
filtered = np.full_like(row, -np.inf)
filtered[top_indices] = row[top_indices]
logits[0] = filtered
row = logits[0]
row = row - np.max(row)
probs = np.exp(row)
probs = probs / probs.sum()
return int(np.random.choice(len(probs), p=probs))
# ============================================================
# Generation function
# ============================================================
def generate(
prompt: str,
max_new_tokens: int = 64,
temperature: float = 0.8,
top_k: int = 50,
) -> str:
"""
Generate text from a prompt using Veylon.
Args:
prompt: Input text
max_new_tokens: Maximum tokens to generate
temperature: Sampling temperature (0.1-2.0)
top_k: Top-K sampling cutoff
Returns:
Generated text
"""
try:
# Encode prompt
tokens = tokenizer.encode(
prompt,
add_bos=True,
add_eos=False,
)
if len(tokens) == 0:
tokens = [tokenizer.bos_id if hasattr(tokenizer, "bos_id") else 1]
tokens = tokens[-CONTEXT:]
# Prefill phase
prompt_ids = np.array([tokens], dtype=np.int32)
logits, cache_k, cache_v = model.generate_step(
prompt_ids,
cache_k=None,
cache_v=None,
cache_pos=0,
)
next_token = sample_from_logits(
np.array(logits[:, -1, :], dtype=np.float32, copy=True),
temperature=temperature,
top_k=top_k,
)
tokens.append(next_token)
# Decoding phase (token-by-token)
if next_token != tokenizer.eos_id and len(tokens) < CONTEXT:
cache_pos = len(prompt_ids[0])
for _ in range(max_new_tokens - 1):
next_input = np.array([[next_token]], dtype=np.int32)
logits, cache_k, cache_v = model.generate_step(
next_input,
cache_k=cache_k,
cache_v=cache_v,
cache_pos=cache_pos,
)
cache_pos += 1
next_token = sample_from_logits(
np.array(logits[:, -1, :], dtype=np.float32, copy=True),
temperature=temperature,
top_k=top_k,
)
tokens.append(next_token)
if next_token == tokenizer.eos_id:
break
if len(tokens) >= CONTEXT:
break
# Decode output
generated_text = tokenizer.decode(tokens)
return generated_text
except Exception as e:
return f"Error: {str(e)}"
# ============================================================
# Gradio UI
# ============================================================
def main():
with gr.Blocks(title="Veylon Alpha") as demo:
gr.Markdown("""
# 🚀 Veylon Alpha Preview - 10M LLM
# Made by Arush Kumar
A student dev!
A small transformer model trained on clean data.
Enter a prompt and watch it generate text.
""")
with gr.Row():
with gr.Column(scale=2):
prompt = gr.Textbox(
label="Prompt",
placeholder="Once upon a time",
lines=3,
value="Once upon a time"
)
with gr.Row():
max_tokens = gr.Slider(
label="Max tokens",
minimum=10,
maximum=256,
value=64,
step=10,
)
temperature = gr.Slider(
label="Temperature",
minimum=0.1,
maximum=2.0,
value=0.8,
step=0.1,
)
top_k = gr.Slider(
label="Top-K",
minimum=1,
maximum=100,
value=50,
step=1,
)
generate_btn = gr.Button("Generate", variant="primary", size="lg")
with gr.Column(scale=1):
info = gr.Markdown(f"""
**Model Info**
- Parameters: {model.count_params():,}
- Context: {CONTEXT} tokens
- Vocab: {vocab_size}
- Architecture: Transformer + GQA
**Tips**
- Higher temp = more creative
- Lower temp = more deterministic
- Top-K = diversity control
""")
output = gr.Textbox(
label="Generated Output",
lines=8,
interactive=False
)
# Connect
generate_btn.click(
fn=generate,
inputs=[prompt, max_tokens, temperature, top_k],
outputs=output,
api_name="generate"
)
demo.launch(share=False, server_name="0.0.0.0", server_port=7860)
if __name__ == "__main__":
main()