from __future__ import annotations import os os.environ["KERAS_BACKEND"] = "jax" import time from pathlib import Path from functools import partial import numpy as np import jax import jax.numpy as jnp import keras import gradio as gr from fastapi import FastAPI import uvicorn 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, ) # ============================================================ # Config # ============================================================ WEIGHTS_PATH = "veylon_final.weights.h5" MODEL_NAME = "Arya" MODEL_TAGLINE = "27M-param transformer, made by Arush Kumar" USER_TAG = "User:" ASSISTANT_TAG = "Assistant:" # ============================================================ # Initialize # ============================================================ keras.mixed_precision.set_global_policy("mixed_bfloat16") print(f"Backend: {keras.backend.backend()}") print(f"JAX devices: {jax.devices()}") 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") 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, ) dummy = np.zeros((1, CONTEXT), dtype=np.int32) _ = model(dummy, training=False) print("Model built successfully") 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") # ============================================================ # JIT-compiled generation steps # ============================================================ JIT_GENERATE = False if JIT_GENERATE: _generate_step_jit = jax.jit(model.generate_step) else: _generate_step_jit = model.generate_step def _run_generate_step(tokens_arr, cache_k, cache_v, cache_pos): return _generate_step_jit( tokens_arr, cache_k=cache_k, cache_v=cache_v, cache_pos=cache_pos, ) # ============================================================ # Sampling # ============================================================ def sample_from_logits( logits: np.ndarray, temperature: float = 0.8, top_k: int = 50, ) -> int: 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)) # ============================================================ # Prompt construction # ============================================================ def build_prompt(message: str, history: list[dict]) -> str: turns = [] for turn in history: role = turn.get("role") content = turn.get("content", "") tag = USER_TAG if role == "user" else ASSISTANT_TAG turns.append(f"{tag} {content}") turns.append(f"{USER_TAG} {message}") turns.append(ASSISTANT_TAG) budget_chars = CONTEXT * 4 kept = [] running_len = 0 for t in reversed(turns): running_len += len(t) + 1 if running_len > budget_chars and kept: break kept.append(t) kept.reverse() return "\n".join(kept) # ============================================================ # Generation # ============================================================ def generate_stream( message: str, history: list[dict], max_new_tokens: int, temperature: float, top_k: int, ): try: prompt = build_prompt(message, history) 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:] prompt_len = len(tokens) prompt_ids = np.array([tokens], dtype=np.int32) t0 = time.time() logits, cache_k, cache_v = _run_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, ) generated_ids = [next_token] partial_text = tokenizer.decode(generated_ids) yield partial_text if next_token != tokenizer.eos_id and prompt_len < CONTEXT: cache_pos = prompt_len for _ in range(int(max_new_tokens) - 1): next_input = np.array([[next_token]], dtype=np.int32) logits, cache_k, cache_v = _run_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, ) if next_token == tokenizer.eos_id: break generated_ids.append(next_token) partial_text = tokenizer.decode(generated_ids) yield partial_text if prompt_len + len(generated_ids) >= CONTEXT: break elapsed = time.time() - t0 n_tok = len(generated_ids) tok_per_sec = n_tok / elapsed if elapsed > 0 else 0.0 print(f"Generated {n_tok} tokens in {elapsed:.2f}s ({tok_per_sec:.1f} tok/s)") except Exception as e: yield f"⚠ Error: {str(e)}" # ============================================================ # Styling — Black Glass & Brushed Metal Aesthetic # ============================================================ CUSTOM_CSS = """ @import url('https://fonts.googleapis.com/css2?family=Inter:wght@300;400;500;600;700&family=JetBrains+Mono:wght@400;500&display=swap'); :root { /* Obsidian & Titanium Palette */ --bg-obsidian: #050507; --bg-metal-dark: #0a0a0c; --bg-metal-light: #141418; /* Glass & Edges */ --glass-bg: rgba(12, 12, 15, 0.65); --glass-edge-top: rgba(255, 255, 255, 0.12); /* Specular highlight */ --glass-edge-bottom: rgba(0, 0, 0, 0.4); --glass-border: rgba(255, 255, 255, 0.06); /* Chrome & Silver Accents */ --chrome-light: #f8fafc; --chrome-mid: #cbd5e1; --chrome-dark: #64748b; /* Text */ --text-primary: #e2e8f0; --text-secondary: #94a3b8; --text-muted: #475569; } * { font-family: 'Inter', -apple-system, BlinkMacSystemFont, sans-serif !important; box-sizing: border-box; } /* Monospace for chat/code/inputs */ .message, .prose, textarea, input, .token-counter, code, pre { font-family: 'JetBrains Mono', ui-monospace, monospace !important; } /* Main Background: Brushed Metal / Anodized Black */ body, .gradio-container { background-color: var(--bg-obsidian) !important; color: var(--text-primary) !important; background-image: /* Subtle top-down lighting */ radial-gradient(ellipse at 50% 0%, rgba(30, 30, 35, 0.4) 0%, transparent 60%), /* Diagonal brushed metal sheen */ linear-gradient(135deg, #08080a 0%, #111115 20%, #0a0a0c 40%, #141418 60%, #08080a 80%, #111115 100%) !important; min-height: 100vh; } /* Header */ #veylon-header { text-align: center; padding: 32px 12px 20px 12px; } #veylon-header h1 { font-size: 2.4rem; font-weight: 700; letter-spacing: -0.5px; /* Polished silver text */ background: linear-gradient(180deg, #ffffff 0%, #94a3b8 60%, #475569 100%); -webkit-background-clip: text; -webkit-text-fill-color: transparent; filter: drop-shadow(0 2px 4px rgba(0,0,0,0.8)); margin-bottom: 8px; } #veylon-header p { color: var(--text-secondary); font-size: 0.85rem; margin: 0; font-weight: 400; letter-spacing: 0.05em; text-transform: uppercase; } /* Glass Panels with Physical Thickness */ .gr-panel, .block, .form, .gap, .wrap { background: var(--glass-bg) !important; border: 1px solid var(--glass-border) !important; /* The magic: top light edge, bottom dark edge for 3D glass effect */ border-top-color: var(--glass-edge-top) !important; border-bottom-color: var(--glass-edge-bottom) !important; border-radius: 16px !important; backdrop-filter: blur(24px) saturate(1.2); -webkit-backdrop-filter: blur(24px) saturate(1.2); box-shadow: 0 12px 40px rgba(0, 0, 0, 0.6), inset 0 1px 0 rgba(255, 255, 255, 0.05), inset 0 -1px 0 rgba(0, 0, 0, 0.2); transition: all 0.3s ease; } /* Chatbot Container (OLED Screen Vibe) */ #chatbot { border-radius: 20px !important; background: rgba(5, 5, 8, 0.85) !important; border: 1px solid rgba(255, 255, 255, 0.04) !important; box-shadow: inset 0 2px 15px rgba(0, 0, 0, 0.8), 0 8px 32px rgba(0, 0, 0, 0.4) !important; overflow: hidden; } #chatbot .form { background: transparent !important; border: none !important; box-shadow: none !important; padding: 10px !important; } /* Chat Bubbles */ .message.user { /* Brushed dark metal for user */ background: linear-gradient(135deg, rgba(35, 35, 40, 0.9) 0%, rgba(20, 20, 25, 0.9) 100%) !important; border: 1px solid rgba(255, 255, 255, 0.08) !important; border-top-color: rgba(255, 255, 255, 0.15) !important; border-radius: 18px 18px 4px 18px !important; color: var(--text-primary) !important; box-shadow: 0 4px 12px rgba(0,0,0,0.5), inset 0 1px 0 rgba(255,255,255,0.05); } .message.bot { /* Smoky glass for bot */ background: rgba(15, 15, 20, 0.6) !important; border: 1px solid rgba(255, 255, 255, 0.04) !important; border-radius: 18px 18px 18px 4px !important; color: var(--text-primary) !important; } .message .message-bubble { padding: 14px 18px !important; line-height: 1.65 !important; font-size: 0.9rem !important; } /* Recessed Inputs (Carved Metal Look) */ input:not([type="range"]), textarea, .gr-box, .gr-text-input { background: rgba(0, 0, 0, 0.5) !important; color: var(--text-primary) !important; border: 1px solid rgba(255, 255, 255, 0.06) !important; border-radius: 12px !important; /* Inset shadow makes it look physically recessed */ box-shadow: inset 0 2px 5px rgba(0, 0, 0, 0.6), inset 0 1px 2px rgba(0,0,0,0.4) !important; transition: all 0.3s ease !important; font-size: 0.9rem !important; } input:not([type="range"]):focus, textarea:focus, .gr-box:focus-within { background: rgba(0, 0, 0, 0.6) !important; border-color: rgba(255, 255, 255, 0.15) !important; box-shadow: inset 0 2px 5px rgba(0, 0, 0, 0.6), 0 0 0 2px rgba(255, 255, 255, 0.05) !important; outline: none !important; } /* Labels */ label span, .form label span { color: var(--text-muted) !important; font-size: 0.7rem !important; text-transform: uppercase !important; letter-spacing: 0.1em !important; font-weight: 600 !important; margin-bottom: 8px !important; } /* Metallic Sliders */ input[type="range"] { -webkit-appearance: none; background: transparent; } input[type="range"]::-webkit-slider-runnable-track { height: 4px; background: linear-gradient(90deg, rgba(255,255,255,0.1), rgba(255,255,255,0.05)); border-radius: 2px; box-shadow: inset 0 1px 2px rgba(0,0,0,0.5); } input[type="range"]::-webkit-slider-thumb { -webkit-appearance: none; height: 18px; width: 18px; border-radius: 50%; /* Chrome sphere */ background: linear-gradient(145deg, #f8fafc 0%, #94a3b8 50%, #475569 100%); margin-top: -7px; box-shadow: 0 2px 6px rgba(0,0,0,0.6), inset 0 1px 0 rgba(255,255,255,0.8); cursor: pointer; transition: transform 0.2s ease; } input[type="range"]::-webkit-slider-thumb:hover { transform: scale(1.15); } /* Polished Chrome Primary Buttons */ button.primary, .gr-button-primary, button[type="submit"] { background: linear-gradient(180deg, #f1f5f9 0%, #94a3b8 40%, #cbd5e1 100%) !important; color: #0f172a !important; border: 1px solid rgba(255, 255, 255, 0.4) !important; border-bottom-color: rgba(0, 0, 0, 0.3) !important; border-radius: 12px !important; font-weight: 700 !important; font-size: 0.85rem !important; letter-spacing: 0.05em; text-transform: uppercase; box-shadow: 0 4px 12px rgba(0, 0, 0, 0.5), inset 0 1px 0 rgba(255, 255, 255, 0.9); text-shadow: 0 1px 0 rgba(255, 255, 255, 0.5); transition: all 0.2s cubic-bezier(0.4, 0, 0.2, 1); padding: 10px 24px !important; } button.primary:hover, .gr-button-primary:hover, button[type="submit"]:hover { background: linear-gradient(180deg, #ffffff 0%, #cbd5e1 40%, #e2e8f0 100%) !important; transform: translateY(-1px); box-shadow: 0 6px 16px rgba(0, 0, 0, 0.6), inset 0 1px 0 rgba(255, 255, 255, 1); } button.primary:active, .gr-button-primary:active { transform: translateY(1px); box-shadow: 0 2px 6px rgba(0, 0, 0, 0.5), inset 0 1px 0 rgba(255, 255, 255, 0.5); } /* Secondary Buttons (Dark Metal) */ button.secondary, .gr-button-secondary, button:not(.primary):not([type="submit"]) { background: linear-gradient(180deg, rgba(30, 30, 35, 0.8) 0%, rgba(15, 15, 20, 0.8) 100%) !important; color: var(--text-secondary) !important; border: 1px solid rgba(255, 255, 255, 0.08) !important; border-top-color: rgba(255, 255, 255, 0.12) !important; border-radius: 12px !important; box-shadow: 0 2px 8px rgba(0,0,0,0.3), inset 0 1px 0 rgba(255,255,255,0.05); transition: all 0.2s ease; } button.secondary:hover, .gr-button-secondary:hover { background: linear-gradient(180deg, rgba(40, 40, 45, 0.9) 0%, rgba(20, 20, 25, 0.9) 100%) !important; color: var(--text-primary) !important; } /* Stats badge */ #stats-badge { text-align: center; color: var(--text-muted); font-size: 0.7rem; padding: 20px 0 4px 0; font-weight: 500; letter-spacing: 0.05em; text-transform: uppercase; } /* Scrollbars */ ::-webkit-scrollbar { width: 6px; height: 6px; } ::-webkit-scrollbar-track { background: transparent; } ::-webkit-scrollbar-thumb { background: rgba(255, 255, 255, 0.1); border-radius: 3px; } ::-webkit-scrollbar-thumb:hover { background: rgba(255, 255, 255, 0.2); } /* Hide default Gradio footer */ footer { display: none !important; } /* Markdown inside chat */ .message .prose h1, .message .prose h2, .message .prose h3 { color: var(--chrome-light) !important; margin-top: 1em !important; } .message .prose p { margin-bottom: 0.8em !important; } .message .prose code { background: rgba(0, 0, 0, 0.5) !important; padding: 2px 6px !important; border-radius: 6px !important; font-size: 0.85em !important; color: var(--chrome-mid) !important; border: 1px solid rgba(255, 255, 255, 0.05) !important; } .message .prose pre { background: rgba(0, 0, 0, 0.6) !important; border: 1px solid rgba(255, 255, 255, 0.05) !important; border-radius: 12px !important; padding: 16px !important; box-shadow: inset 0 2px 8px rgba(0,0,0,0.4); } """ THEME = gr.themes.Base( primary_hue="slate", secondary_hue="zinc", neutral_hue="zinc", font=[gr.themes.GoogleFont("Inter"), "ui-sans-serif", "system-ui", "sans-serif"], ).set( body_background_fill="#050507", block_background_fill="rgba(12, 12, 15, 0.65)", border_color_primary="rgba(255, 255, 255, 0.06)", ) # ============================================================ # Build the Gradio Blocks app # ============================================================ def build_demo() -> gr.Blocks: with gr.Blocks(title=MODEL_NAME, css=CUSTOM_CSS, theme=THEME) as demo: gr.HTML(f"""
{MODEL_TAGLINE} · {model.count_params():,} params · {CONTEXT} ctx