Veylon / veylon_model.py
Arush kumar
Update veylon_model.py
d4f1a09
Raw
History Blame Contribute Delete
22.4 kB
from __future__ import annotations
import numpy as np
import keras
from keras import layers, ops
import jax
from veylon_attention import flash_splash_attention, decode_swa
try:
from config import (
CONTEXT,
vocab_size as Vocab_size,
D_MODEL,
numberoflayers,
numberofheads,
d_Latent,
ffn_mult,
swa_window,
num_kv_heads,
use_moe,
moe_num_experts,
moe_top_k,
USE_REMAT,
MAX_GEN_TOKENS,
)
except Exception:
CONTEXT = 2048
Vocab_size = 32000
D_MODEL = 512
numberoflayers = 8
numberofheads = 8
d_Latent = 128
ffn_mult = 3.5
swa_window = 1024
num_kv_heads = 2
use_moe = False
moe_num_experts = 8
moe_top_k = 2
USE_REMAT = True
MAX_GEN_TOKENS = 256
# NOTE: mixed precision policy is set by train.py / finetune.py before model
# creation. Do NOT set it here β€” importing this module must not have side
# effects.
# ─────────────────────────────────────────────────────────────────────────────
# Layers
# ─────────────────────────────────────────────────────────────────────────────
@keras.saving.register_keras_serializable()
class RMSNorm(layers.Layer):
def __init__(self, epsilon=1e-5, **kwargs):
super().__init__(**kwargs)
self.epsilon = epsilon
def build(self, input_shape):
self.weight = self.add_weight(shape=(input_shape[-1],), initializer='ones', name='gamma')
def call(self, x):
x_fp32 = ops.cast(x, 'float32')
rms = ops.sqrt(ops.mean(ops.square(x_fp32), axis=-1, keepdims=True) + self.epsilon)
out = x_fp32 / rms
return ops.cast(out, x.dtype) * self.weight
def get_config(self):
cfg = super().get_config()
cfg.update({'epsilon': self.epsilon})
return cfg
@keras.saving.register_keras_serializable()
class RotaryEmbedding(layers.Layer):
"""
RoPE with `max_seq_len + gen_headroom` so a full-context prompt can still
decode up to `gen_headroom` tokens without overflowing the cos/sin table.
`max_seq_len` is preserved in the config so checkpoints remain compatible β€”
only the (trainable=False) table shapes grow.
"""
def __init__(self, max_seq_len, dim, theta=10000.0, gen_headroom=0, **kwargs):
super().__init__(**kwargs)
if dim % 2 != 0:
raise ValueError('RotaryEmbedding dim must be even.')
self.max_seq_len = max_seq_len
self.dim = dim
self.theta = theta
self.gen_headroom = int(gen_headroom)
self.table_size = max_seq_len + self.gen_headroom
def build(self, input_shape):
half = self.dim // 2
inv_freq = 1.0 / (self.theta ** (np.arange(0, self.dim, 2).astype(np.float32) / self.dim))
positions = np.arange(self.table_size, dtype=np.float32)
freqs = positions[:, None] * inv_freq[None, :]
self.cos = self.add_weight(
shape=(self.table_size, half),
initializer=keras.initializers.Constant(np.cos(freqs)),
trainable=False, dtype='float32', name='cos_table',
)
self.sin = self.add_weight(
shape=(self.table_size, half),
initializer=keras.initializers.Constant(np.sin(freqs)),
trainable=False, dtype='float32', name='sin_table',
)
def call(self, x, offset=0):
seq_len = x.shape[1]
if seq_len is None:
seq_len = ops.shape(x)[1]
half = self.dim // 2
if offset + seq_len > self.table_size:
raise ValueError(
f'RoPE table too small: offset={offset}, seq_len={seq_len}, '
f'max={self.table_size} (CONTEXT={self.max_seq_len} + gen_headroom={self.gen_headroom})'
)
cos = self.cos[offset:offset + seq_len, :half]
sin = self.sin[offset:offset + seq_len, :half]
cos = ops.cast(ops.reshape(cos, (1, seq_len, 1, half)), x.dtype)
sin = ops.cast(ops.reshape(sin, (1, seq_len, 1, half)), x.dtype)
x1 = x[..., :half]
x2 = x[..., half:]
return ops.concatenate([x1 * cos - x2 * sin, x1 * sin + x2 * cos], axis=-1)
def get_config(self):
cfg = super().get_config()
cfg.update({
'max_seq_len': self.max_seq_len,
'dim': self.dim,
'theta': self.theta,
'gen_headroom': self.gen_headroom,
})
return cfg
@keras.saving.register_keras_serializable()
class SwiGLUFFN(layers.Layer):
def __init__(self, d_model, hidden_mult=3.5, **kwargs):
super().__init__(**kwargs)
self.d_model_arg = d_model
self.hidden_mult = hidden_mult
self.hidden_dim = int(d_model * hidden_mult * 2 / 3)
self.hidden_dim = ((self.hidden_dim + 63) // 64) * 64
def build(self, input_shape):
d_model = input_shape[-1]
self.gate_up_proj = self.add_weight(shape=(d_model, 2 * self.hidden_dim), initializer='glorot_uniform', name='gate_up_proj')
self.down_proj = self.add_weight(shape=(self.hidden_dim, d_model), initializer='glorot_uniform', name='down_proj')
def call(self, x, training=False):
gate_up = ops.matmul(x, self.gate_up_proj)
gate, up = ops.split(gate_up, 2, axis=-1)
return ops.matmul(ops.silu(gate) * up, self.down_proj)
def get_config(self):
cfg = super().get_config()
cfg.update({'d_model': self.d_model_arg, 'hidden_mult': self.hidden_mult})
return cfg
@keras.saving.register_keras_serializable()
class MoE_FFN(layers.Layer):
"""
Mixture-of-Experts FFN (top-k routing, masked dispatch).
For each routing slot k in [0, top_k) we run each expert once over the
full batch, zeroing inputs the expert was not selected for, then accumulate
the gated expert output. The load-balancing loss uses top-1 argmax routing
for the `f` term (standard Switch-Transformer formulation); the previous
code used top-k membership which over-counted and produced the wrong loss.
"""
def __init__(self, d_model, num_experts=8, top_k=2, hidden_mult=3.5, **kwargs):
super().__init__(**kwargs)
self.d_model_arg = d_model
self.num_experts = num_experts
self.top_k = top_k
self.hidden_mult = hidden_mult
self.experts = [SwiGLUFFN(d_model, hidden_mult) for _ in range(num_experts)]
self.router = layers.Dense(num_experts, use_bias=False)
def _load_balancing_loss(self, router_logits, top_idx):
# Standard Switch-Transformer aux loss. `f` uses the argmax routing
# decision (top_idx[..., 0]); `p` uses the router softmax mean.
router_probs = ops.softmax(router_logits, axis=-1)
mask = ops.one_hot(top_idx[..., 0], self.num_experts) # [..., num_experts]
f = ops.mean(ops.cast(mask, 'float32'), axis=tuple(range(mask.ndim - 1)))
p = ops.mean(router_probs, axis=tuple(range(router_probs.ndim - 1)))
return self.num_experts * ops.sum(f * p)
def call(self, x, training=False):
# x: [B, S, D]
B = ops.shape(x)[0]
S = ops.shape(x)[1]
D = x.shape[-1]
router_logits = self.router(x) # [B, S, E]
top_logits, top_idx = ops.top_k(router_logits, self.top_k) # [B, S, k]
top_weights = ops.softmax(top_logits, axis=-1) # [B, S, k]
if training:
self.add_loss(self._load_balancing_loss(router_logits, top_idx))
out = ops.zeros_like(x)
for kk in range(self.top_k):
idx_k = top_idx[..., kk] # [B, S]
w_k = top_weights[..., kk][..., None] # [B, S, 1]
for e in range(self.num_experts):
mask_e = ops.cast(idx_k == e, x.dtype)[..., None] # [B, S, 1]
# mask_e is always a tensor; skip the useless ops.is_tensor guard
# (it never triggered β€” removed FIX BUG11)
# Only run the expert where it is actually selected.
x_e = x * mask_e
expert_out = self.experts[e](x_e, training=training)
out = out + expert_out * w_k * mask_e
return out
def get_config(self):
cfg = super().get_config()
cfg.update({
'd_model': self.d_model_arg, 'num_experts': self.num_experts,
'top_k': self.top_k, 'hidden_mult': self.hidden_mult,
})
return cfg
@keras.saving.register_keras_serializable()
class MLAttention(layers.Layer):
def __init__(self, d_model, n_heads, d_latent, max_seq_len, num_kv_heads=2,
swa_window=1024, attn_dropout=0.0, gen_headroom=0, **kwargs):
super().__init__(**kwargs)
if d_model % n_heads != 0:
raise ValueError('d_model must be divisible by n_heads')
if n_heads % num_kv_heads != 0:
raise ValueError('n_heads must be divisible by num_kv_heads')
self.d_model = d_model
self.n_heads = n_heads
self.num_kv_heads = num_kv_heads
self.group_size = n_heads // num_kv_heads
self.d_head = d_model // n_heads
self.d_latent = d_latent
self.max_seq_len = max_seq_len
self.swa_window = swa_window
self.gen_headroom = int(gen_headroom)
self.dropout = layers.Dropout(attn_dropout)
self.rope = RotaryEmbedding(max_seq_len, d_model // n_heads, gen_headroom=int(gen_headroom)) # FIX BUG10: must be in __init__ for Keras tracking
def build(self, input_shape):
self.W_qc = self.add_weight(shape=(self.d_model, self.d_model + self.d_latent), initializer='glorot_uniform', name='W_qc')
self.W_kv = self.add_weight(shape=(self.d_latent, self.num_kv_heads * 2 * self.d_head), initializer='glorot_uniform', name='W_kv')
self.W_o = self.add_weight(shape=(self.d_model, self.d_model), initializer='glorot_uniform', name='Wo')
def _project_kv(self, c):
kv = ops.matmul(c, self.W_kv)
return ops.split(kv, 2, axis=-1)
def call(self, x, training=False):
B = ops.shape(x)[0]
S = ops.shape(x)[1]
qc = ops.matmul(x, self.W_qc)
q_proj, c = ops.split(qc, [self.d_model], axis=-1)
q = ops.reshape(q_proj, (B, S, self.n_heads, self.d_head))
q = self.rope(q, offset=0)
q = ops.transpose(q, (0, 2, 1, 3))
k, v = self._project_kv(c)
k = ops.reshape(k, (B, S, self.num_kv_heads, self.d_head))
v = ops.reshape(v, (B, S, self.num_kv_heads, self.d_head))
k = self.rope(k, offset=0)
k = ops.transpose(k, (0, 2, 1, 3))
v = ops.transpose(v, (0, 2, 1, 3))
# Attention remat is REQUIRED during training: it is the per-block
# jax.checkpoint(process_block) that stops XLA from staging the hidden
# [S, B, Hkv, W, D] buffer (β‰ˆ8.6 GB here) inside the fori_loop. Block-
# level remat does NOT prevent that forward-time allocation, so this
# must follow `training` regardless of config.USE_REMAT.
out = flash_splash_attention(
q, k, v,
window_size=min(self.swa_window, self.max_seq_len),
backend=jax.default_backend(),
use_gqa=True,
use_remat=training,
)
out = ops.transpose(out, (0, 2, 1, 3))
out = ops.reshape(out, (B, S, self.d_model))
out = self.dropout(out, training=training)
return ops.matmul(out, self.W_o)
def generate_step(self, x, cache_k=None, cache_v=None, cache_pos=0):
B = ops.shape(x)[0]
S = ops.shape(x)[1]
qc = ops.matmul(x, self.W_qc)
q_proj, c = ops.split(qc, [self.d_model], axis=-1)
q = ops.reshape(q_proj, (B, S, self.n_heads, self.d_head))
q = self.rope(q, offset=cache_pos)
q = ops.transpose(q, (0, 2, 1, 3))
k, v = self._project_kv(c)
k = ops.reshape(k, (B, S, self.num_kv_heads, self.d_head))
v = ops.reshape(v, (B, S, self.num_kv_heads, self.d_head))
k = self.rope(k, offset=cache_pos)
k = ops.transpose(k, (0, 2, 1, 3))
v = ops.transpose(v, (0, 2, 1, 3))
if cache_k is None:
# Prefill: no remat (inference only), no window trim needed.
out = flash_splash_attention(
q, k, v,
window_size=min(self.swa_window, self.max_seq_len),
backend=jax.default_backend(),
use_gqa=True,
use_remat=False,
)
new_k = k[:, :, -self.swa_window:, :]
new_v = v[:, :, -self.swa_window:, :]
else:
if S != 1:
raise ValueError(f'generate_step with cache expects S=1, got S={S}')
k = ops.concatenate([cache_k, k], axis=2)
v = ops.concatenate([cache_v, v], axis=2)
k = k[:, :, -self.swa_window:, :]
v = v[:, :, -self.swa_window:, :]
out = decode_swa(q, k, v)
new_k = k
new_v = v
out = ops.transpose(out, (0, 2, 1, 3))
out = ops.reshape(out, (B, S, self.d_model))
out = ops.matmul(out, self.W_o)
return out, new_k, new_v
def get_config(self):
cfg = super().get_config()
cfg.update({
'd_model': self.d_model, 'n_heads': self.n_heads,
'num_kv_heads': self.num_kv_heads, 'd_latent': self.d_latent,
'max_seq_len': self.max_seq_len, 'swa_window': self.swa_window,
'attn_dropout': self.dropout.rate, 'gen_headroom': self.gen_headroom,
})
return cfg
@keras.saving.register_keras_serializable()
class TransformerBlock(layers.Layer):
def __init__(self, d_model, n_heads, d_latent, ffn_layer, max_seq_len,
num_kv_heads=2, swa_window=1024, use_remat=True, gen_headroom=0, **kwargs):
super().__init__(**kwargs)
self.d_model = d_model
self.n_heads = n_heads
self.d_latent = d_latent
self.max_seq_len = max_seq_len
self.num_kv_heads = num_kv_heads
self.swa_window = swa_window
self.use_remat = use_remat
self.gen_headroom = int(gen_headroom)
self.ffn = keras.saving.deserialize_keras_object(ffn_layer) if isinstance(ffn_layer, dict) else ffn_layer
self.norm1 = RMSNorm()
self.norm2 = RMSNorm()
self.attn = MLAttention(
d_model, n_heads, d_latent, max_seq_len,
num_kv_heads=num_kv_heads, swa_window=swa_window, gen_headroom=self.gen_headroom,
)
def call(self, x, training=False):
# Block-level activation/gradient checkpointing (optional, via
# config.USE_REMAT). NOTE: this is independent of attention's own
# remat, which MUST stay on whenever training=True β€” that's what stops
# XLA staging the hidden [S,B,Hkv,W,D] buffer (the OOM we hit).
if training and self.use_remat:
def _fwd(x_in):
attn_out = self.attn(self.norm1(x_in), training=True)
ffn_out = self.ffn(self.norm2(x_in), training=True)
return attn_out + ffn_out
return x + jax.checkpoint(_fwd)(x)
a = self.attn(self.norm1(x), training=training)
f = self.ffn(self.norm2(x), training=training)
return x + a + f
def generate_step(self, x, cache_k=None, cache_v=None, cache_pos=0):
# Must mirror call()'s PARALLEL residual structure exactly:
# attn and ffn both read from norm(x) computed on the SAME
# pre-block x, then both get added to that same original x.
# Previously this was written as a sequential block (x=x+attn;
# x=x+ffn(norm2(x))) β€” ffn ended up conditioned on norm2(x+attn_out)
# instead of norm2(x), an input distribution the FFN weights were
# never trained on. That silently produced a different model at
# generation time than at training time: teacher-forced loss (via
# call()) never exercises this path and looked fine, while actual
# autoregressive generation compounded the mismatch every step.
attn_out, nck, ncv = self.attn.generate_step(
self.norm1(x), cache_k=cache_k, cache_v=cache_v, cache_pos=cache_pos,
)
ffn_out = self.ffn(self.norm2(x), training=False)
x = x + attn_out + ffn_out
return x, nck, ncv
def get_config(self):
cfg = super().get_config()
cfg.update({
'd_model': self.d_model, 'n_heads': self.n_heads, 'd_latent': self.d_latent,
'ffn_layer': keras.saving.serialize_keras_object(self.ffn),
'max_seq_len': self.max_seq_len, 'num_kv_heads': self.num_kv_heads,
'swa_window': self.swa_window, 'use_remat': self.use_remat,
'gen_headroom': self.gen_headroom,
})
return cfg
@keras.saving.register_keras_serializable()
class VeylonModel(keras.Model):
def __init__(self, vocab_size, d_model, n_layers, n_heads, d_latent, ffn_mult,
max_seq_len, use_moe=False, moe_num_experts=8, moe_top_k=2,
num_kv_heads=2, swa_window=1024, use_remat=True, gen_headroom=0, **kwargs):
super().__init__(**kwargs)
self.vocab_size = vocab_size
self.d_model = d_model
self.n_layers = n_layers
self.n_heads = n_heads
self.d_latent = d_latent
self.ffn_mult = ffn_mult
self.max_seq_len = max_seq_len
self.use_moe = use_moe
self.moe_num_experts = moe_num_experts
self.moe_top_k = moe_top_k
self.num_kv_heads = num_kv_heads
self.swa_window = swa_window
self.use_remat = use_remat
self.gen_headroom = int(gen_headroom)
self.embedding = layers.Embedding(vocab_size, d_model, name='token_embedding')
self.blocks = []
for i in range(n_layers):
ffn = (MoE_FFN(d_model, moe_num_experts, moe_top_k, ffn_mult)
if use_moe else SwiGLUFFN(d_model, ffn_mult))
self.blocks.append(TransformerBlock(
d_model, n_heads, d_latent, ffn, max_seq_len,
num_kv_heads=num_kv_heads, swa_window=swa_window,
use_remat=use_remat, gen_headroom=self.gen_headroom, name=f'block_{i}',
))
self.norm = RMSNorm()
def call(self, inputs, training=False):
x = self.embedding(inputs)
for block in self.blocks:
x = block(x, training=training)
x = self.norm(x)
embedding_weights = self.embedding.embeddings
logits = ops.matmul(x, ops.transpose(embedding_weights))
return ops.cast(logits, 'float32')
def generate_step(self, inputs, cache_k=None, cache_v=None, cache_pos=0):
x = self.embedding(inputs)
new_cache_k = []
new_cache_v = []
if cache_k is None:
cache_k = [None] * len(self.blocks)
cache_v = [None] * len(self.blocks)
for i, block in enumerate(self.blocks):
x, nck, ncv = block.generate_step(
x, cache_k=cache_k[i], cache_v=cache_v[i], cache_pos=cache_pos,
)
new_cache_k.append(nck)
new_cache_v.append(ncv)
x = self.norm(x)
embedding_weights = self.embedding.embeddings
logits = ops.matmul(x, ops.transpose(embedding_weights))
logits = ops.cast(logits, 'float32')
return logits, new_cache_k, new_cache_v
def get_config(self):
cfg = super().get_config()
cfg.update({
'vocab_size': self.vocab_size, 'd_model': self.d_model,
'n_layers': self.n_layers, 'n_heads': self.n_heads,
'd_latent': self.d_latent, 'ffn_mult': self.ffn_mult,
'max_seq_len': self.max_seq_len, 'use_moe': self.use_moe,
'moe_num_experts': self.moe_num_experts, 'moe_top_k': self.moe_top_k,
'num_kv_heads': self.num_kv_heads, 'swa_window': self.swa_window,
'use_remat': self.use_remat, 'gen_headroom': self.gen_headroom,
})
return cfg
# ─────────────────────────────────────────────────────────────────────────────
# Factory
# ─────────────────────────────────────────────────────────────────────────────
def 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=use_moe,
moe_num_experts=moe_num_experts, moe_top_k=moe_top_k, num_kv_heads=num_kv_heads,
swa_window=swa_window, use_remat=USE_REMAT, gen_headroom=MAX_GEN_TOKENS,
):
"""
Build a VeylonModel.
`gen_headroom` sizes the RoPE table to (max_seq_len + gen_headroom) so that
a full-length prompt can still generate `gen_headroom` tokens during
inference without "RoPE table too small" errors. Defaults to MAX_GEN_TOKENS.
"""
return VeylonModel(
vocab_size=vocab_size, d_model=d_model, n_layers=n_layers, n_heads=n_heads,
d_latent=d_latent, ffn_mult=ffn_mult, max_seq_len=max_seq_len, use_moe=use_moe,
moe_num_experts=moe_num_experts, moe_top_k=moe_top_k, num_kv_heads=num_kv_heads,
swa_window=swa_window, use_remat=use_remat, gen_headroom=gen_headroom,
)