Spaces:
Sleeping
Sleeping
| 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) |