from __future__ import annotations import os os.environ["KERAS_BACKEND"] = "jax" import numpy as np import jax import keras 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, ) # ============================================================ # Runtime info # ============================================================ print(f"Backend: {keras.backend.backend()}") print(f"JAX devices: {jax.devices()}") keras.mixed_precision.set_global_policy("mixed_bfloat16") # ============================================================ # Load tokenizer # ============================================================ tokenizer = TokenizerWrapper("tokenizer.model") assert tokenizer.vocab_size == vocab_size, ( f"Tokenizer vocab ({tokenizer.vocab_size}) " f"!= config vocab ({vocab_size})" ) print(f"Tokenizer vocab size: {tokenizer.vocab_size}") # ============================================================ # Build model (must exactly match training) # ============================================================ 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 with EXACT training/inference shape dummy = np.zeros((1, CONTEXT), dtype=np.int32) _ = model(dummy, training=False) print("✓ Model built successfully") # ============================================================ # Load weights # ============================================================ WEIGHTS_PATH = "veylon_final.weights.h5" print(f"Loading weights from: {WEIGHTS_PATH}") model.load_weights(WEIGHTS_PATH) print("✓ Weights loaded successfully") # ============================================================ # Sampling settings # ============================================================ MAX_NEW_TOKENS = 64 TEMPERATURE = 0.8 TOP_K = 50 def sample_from_logits( logits: np.ndarray, temperature: float = 0.8, top_k: int = 50, ) -> int: """ NumPy-only sampling to avoid JAX/readonly array issues. """ 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 loop # ============================================================ while True: prompt = input("\nEnter your prompt (or 'exit'): ").strip() if prompt.lower() in {"exit", "quit"}: break 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:] print("\nGenerating...\n") # Prompt prefill (one-time) 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) if next_token != tokenizer.eos_id and len(tokens) < CONTEXT: # After prefill, we are decoding token-by-token. 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: print("\n[Context limit reached]") break generated_text = tokenizer.decode(tokens) print("\n" + "=" * 60) print("Veylon Alpha") print("=" * 60) print(generated_text) print("=" * 60)