Veylon / inference.py
Arush kumar
Upload 14 files
54ad1e5
Raw
History Blame Contribute Delete
5.19 kB
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)