Pragya / veylon_model.py
Arush kumar
Upload 14 files
54ad1e5
Raw History Blame
15.9 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,
)
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
keras.mixed_precision.set_global_policy('mixed_bfloat16')
@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):
def __init__(self, max_seq_len, dim, theta=10000.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
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.max_seq_len, dtype=np.float32)
freqs = positions[:, None] * inv_freq[None, :]
self.cos = self.add_weight(shape=(self.max_seq_len, half), initializer=keras.initializers.Constant(np.cos(freqs)), trainable=False, dtype='float32', name='cos_table')
self.sin = self.add_weight(shape=(self.max_seq_len, 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.max_seq_len:
raise ValueError(f'RoPE table too small: offset={offset}, seq_len={seq_len}, max={self.max_seq_len}')
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})
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):
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_k_indices):
router_probs = ops.softmax(router_logits, axis=-1)
mask = ops.one_hot(top_k_indices, self.num_experts)
mask = ops.sum(mask, axis=2)
f = ops.mean(mask, axis=(0, 1))
p = ops.mean(router_probs, axis=(0, 1))
return self.num_experts * ops.sum(f * p)
def call(self, x, training=False):
router_logits = self.router(x)
top_logits, top_idx = ops.top_k(router_logits, self.top_k)
top_weights = ops.softmax(top_logits, axis=-1)
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 = top_idx[..., kk]
w = top_weights[..., kk]
for e in range(self.num_experts):
mask = ops.cast(idx == e, x.dtype)
expert_out = self.experts[e](x * mask[..., None], training=training)
out += expert_out * w[..., None] * mask[..., None]
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, **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.dropout = layers.Dropout(attn_dropout)
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')
self.rope = RotaryEmbedding(self.max_seq_len, self.d_head)
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))
out = flash_splash_attention(
q,
k,
v,
window_size=min(self.swa_window, self.max_seq_len),
backend=jax.default_backend(),
use_gqa=True,
)
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))
# Prefill path: no previous cache
if cache_k is None:
out = flash_splash_attention(
q,
k,
v,
window_size=min(self.swa_window, self.max_seq_len),
backend=jax.default_backend(),
use_gqa=True,
)
new_k = k[:, :, -self.swa_window :, :]
new_v = v[:, :, -self.swa_window :, :]
else:
# Decode path: one-token generation only
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})
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, **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.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)
def call(self, x, training=False):
x = x + self.attn(self.norm1(x), training=training)
x = x + self.ffn(self.norm2(x), training=training)
return x
def generate_step(self, x, cache_k=None, cache_v=None, cache_pos=0):
attn_out, nck, ncv = self.attn.generate_step(
self.norm1(x),
cache_k=cache_k,
cache_v=cache_v,
cache_pos=cache_pos,
)
x = x + attn_out
x = x + self.ffn(self.norm2(x), training=False)
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})
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, **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.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, 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.weights[0]
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.weights[0]
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})
return cfg
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):
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)