Spaces:
Sleeping
Sleeping
Arush kumar commited on
Commit ·
54ad1e5
1
Parent(s): a098b7a
Upload 14 files
Browse files- check_vocab.py +4 -0
- config.py +25 -0
- corpus.txt +1 -0
- inference.py +193 -0
- inspect_weights.py +7 -0
- model_build.py +56 -0
- token_train.py +20 -0
- tokenizer.model +3 -0
- tokenizer.vocab +0 -0
- train.py +439 -0
- train2.py +547 -0
- veylon_attention.py +640 -0
- veylon_model.py +383 -0
check_vocab.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
|
| 3 |
+
v = json.load(open('vocab.json'))
|
| 4 |
+
print(f'Current vocab size: {len(v["char_to_idx"])}')
|
config.py
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
EPOCHS = 20
|
| 4 |
+
CONTEXT = 1024
|
| 5 |
+
data_path = "./training_data"
|
| 6 |
+
batch_size = 8
|
| 7 |
+
learning_rate = 1e-4
|
| 8 |
+
weight_decay = 0.01
|
| 9 |
+
|
| 10 |
+
vocab_size = 8000
|
| 11 |
+
D_MODEL = 256
|
| 12 |
+
numberoflayers = 8
|
| 13 |
+
numberofheads = 4
|
| 14 |
+
num_kv_heads = 2
|
| 15 |
+
d_Latent = 64
|
| 16 |
+
ffn_mult = 3.5
|
| 17 |
+
swa_window = 512
|
| 18 |
+
use_moe = False
|
| 19 |
+
moe_num_experts = 8
|
| 20 |
+
moe_top_k = 2
|
| 21 |
+
|
| 22 |
+
GLOBAL_DTYPE = "bfloat16" # "float32", "float16", "bfloat16"
|
| 23 |
+
MATMUL_PRECISION = "high" # "default", "high", "highest"
|
| 24 |
+
USE_FLASH_ATTENTION = False
|
| 25 |
+
USE_SPLASH_ATTENTION = True
|
corpus.txt
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs. Advanced reasoning involves breaking down problems into steps, critiquing solutions, and iteratively improving answers. Chain-of-thought improves logical consistency. Self-critique helps models identify flaws in their own reasoning. Recursive refinement (TRM) leads to better final outputs.
|
inference.py
ADDED
|
@@ -0,0 +1,193 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
os.environ["KERAS_BACKEND"] = "jax"
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import jax
|
| 8 |
+
import keras
|
| 9 |
+
|
| 10 |
+
from veylon_model import create_llm
|
| 11 |
+
from tokenizer import TokenizerWrapper
|
| 12 |
+
|
| 13 |
+
from config import (
|
| 14 |
+
CONTEXT,
|
| 15 |
+
vocab_size,
|
| 16 |
+
D_MODEL,
|
| 17 |
+
numberoflayers,
|
| 18 |
+
numberofheads,
|
| 19 |
+
d_Latent,
|
| 20 |
+
ffn_mult,
|
| 21 |
+
num_kv_heads,
|
| 22 |
+
swa_window,
|
| 23 |
+
)
|
| 24 |
+
|
| 25 |
+
# ============================================================
|
| 26 |
+
# Runtime info
|
| 27 |
+
# ============================================================
|
| 28 |
+
|
| 29 |
+
print(f"Backend: {keras.backend.backend()}")
|
| 30 |
+
print(f"JAX devices: {jax.devices()}")
|
| 31 |
+
|
| 32 |
+
keras.mixed_precision.set_global_policy("mixed_bfloat16")
|
| 33 |
+
|
| 34 |
+
# ============================================================
|
| 35 |
+
# Load tokenizer
|
| 36 |
+
# ============================================================
|
| 37 |
+
|
| 38 |
+
tokenizer = TokenizerWrapper("tokenizer.model")
|
| 39 |
+
|
| 40 |
+
assert tokenizer.vocab_size == vocab_size, (
|
| 41 |
+
f"Tokenizer vocab ({tokenizer.vocab_size}) "
|
| 42 |
+
f"!= config vocab ({vocab_size})"
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
print(f"Tokenizer vocab size: {tokenizer.vocab_size}")
|
| 46 |
+
|
| 47 |
+
# ============================================================
|
| 48 |
+
# Build model (must exactly match training)
|
| 49 |
+
# ============================================================
|
| 50 |
+
|
| 51 |
+
print("Building model...")
|
| 52 |
+
|
| 53 |
+
model = create_llm(
|
| 54 |
+
vocab_size=vocab_size,
|
| 55 |
+
d_model=D_MODEL,
|
| 56 |
+
n_layers=numberoflayers,
|
| 57 |
+
n_heads=numberofheads,
|
| 58 |
+
d_latent=d_Latent,
|
| 59 |
+
ffn_mult=ffn_mult,
|
| 60 |
+
max_seq_len=CONTEXT,
|
| 61 |
+
use_moe=False,
|
| 62 |
+
num_kv_heads=num_kv_heads,
|
| 63 |
+
swa_window=swa_window,
|
| 64 |
+
)
|
| 65 |
+
|
| 66 |
+
# Warmup with EXACT training/inference shape
|
| 67 |
+
dummy = np.zeros((1, CONTEXT), dtype=np.int32)
|
| 68 |
+
_ = model(dummy, training=False)
|
| 69 |
+
|
| 70 |
+
print("✓ Model built successfully")
|
| 71 |
+
|
| 72 |
+
# ============================================================
|
| 73 |
+
# Load weights
|
| 74 |
+
# ============================================================
|
| 75 |
+
|
| 76 |
+
WEIGHTS_PATH = "veylon_final.weights.h5"
|
| 77 |
+
print(f"Loading weights from: {WEIGHTS_PATH}")
|
| 78 |
+
model.load_weights(WEIGHTS_PATH)
|
| 79 |
+
print("✓ Weights loaded successfully")
|
| 80 |
+
|
| 81 |
+
# ============================================================
|
| 82 |
+
# Sampling settings
|
| 83 |
+
# ============================================================
|
| 84 |
+
|
| 85 |
+
MAX_NEW_TOKENS = 64
|
| 86 |
+
TEMPERATURE = 0.8
|
| 87 |
+
TOP_K = 50
|
| 88 |
+
|
| 89 |
+
def sample_from_logits(
|
| 90 |
+
logits: np.ndarray,
|
| 91 |
+
temperature: float = 0.8,
|
| 92 |
+
top_k: int = 50,
|
| 93 |
+
) -> int:
|
| 94 |
+
"""
|
| 95 |
+
NumPy-only sampling to avoid JAX/readonly array issues.
|
| 96 |
+
"""
|
| 97 |
+
logits = np.array(logits, dtype=np.float32, copy=True)
|
| 98 |
+
|
| 99 |
+
if temperature > 0:
|
| 100 |
+
logits = logits / float(max(temperature, 1e-8))
|
| 101 |
+
|
| 102 |
+
if top_k > 0:
|
| 103 |
+
k = min(int(top_k), logits.shape[-1])
|
| 104 |
+
row = logits[0]
|
| 105 |
+
top_indices = np.argpartition(row, -k)[-k:]
|
| 106 |
+
|
| 107 |
+
filtered = np.full_like(row, -np.inf)
|
| 108 |
+
filtered[top_indices] = row[top_indices]
|
| 109 |
+
logits[0] = filtered
|
| 110 |
+
|
| 111 |
+
row = logits[0]
|
| 112 |
+
row = row - np.max(row)
|
| 113 |
+
probs = np.exp(row)
|
| 114 |
+
probs = probs / probs.sum()
|
| 115 |
+
|
| 116 |
+
return int(np.random.choice(len(probs), p=probs))
|
| 117 |
+
|
| 118 |
+
# ============================================================
|
| 119 |
+
# Generation loop
|
| 120 |
+
# ============================================================
|
| 121 |
+
|
| 122 |
+
while True:
|
| 123 |
+
prompt = input("\nEnter your prompt (or 'exit'): ").strip()
|
| 124 |
+
|
| 125 |
+
if prompt.lower() in {"exit", "quit"}:
|
| 126 |
+
break
|
| 127 |
+
|
| 128 |
+
tokens = tokenizer.encode(
|
| 129 |
+
prompt,
|
| 130 |
+
add_bos=True,
|
| 131 |
+
add_eos=False,
|
| 132 |
+
)
|
| 133 |
+
|
| 134 |
+
if len(tokens) == 0:
|
| 135 |
+
tokens = [tokenizer.bos_id if hasattr(tokenizer, "bos_id") else 1]
|
| 136 |
+
|
| 137 |
+
tokens = tokens[-CONTEXT:]
|
| 138 |
+
|
| 139 |
+
print("\nGenerating...\n")
|
| 140 |
+
|
| 141 |
+
# Prompt prefill (one-time)
|
| 142 |
+
prompt_ids = np.array([tokens], dtype=np.int32)
|
| 143 |
+
logits, cache_k, cache_v = model.generate_step(
|
| 144 |
+
prompt_ids,
|
| 145 |
+
cache_k=None,
|
| 146 |
+
cache_v=None,
|
| 147 |
+
cache_pos=0,
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
next_token = sample_from_logits(
|
| 151 |
+
np.array(logits[:, -1, :], dtype=np.float32, copy=True),
|
| 152 |
+
temperature=TEMPERATURE,
|
| 153 |
+
top_k=TOP_K,
|
| 154 |
+
)
|
| 155 |
+
tokens.append(next_token)
|
| 156 |
+
|
| 157 |
+
if next_token != tokenizer.eos_id and len(tokens) < CONTEXT:
|
| 158 |
+
# After prefill, we are decoding token-by-token.
|
| 159 |
+
cache_pos = len(prompt_ids[0])
|
| 160 |
+
|
| 161 |
+
for _ in range(MAX_NEW_TOKENS - 1):
|
| 162 |
+
next_input = np.array([[next_token]], dtype=np.int32)
|
| 163 |
+
|
| 164 |
+
logits, cache_k, cache_v = model.generate_step(
|
| 165 |
+
next_input,
|
| 166 |
+
cache_k=cache_k,
|
| 167 |
+
cache_v=cache_v,
|
| 168 |
+
cache_pos=cache_pos,
|
| 169 |
+
)
|
| 170 |
+
|
| 171 |
+
cache_pos += 1
|
| 172 |
+
|
| 173 |
+
next_token = sample_from_logits(
|
| 174 |
+
np.array(logits[:, -1, :], dtype=np.float32, copy=True),
|
| 175 |
+
temperature=TEMPERATURE,
|
| 176 |
+
top_k=TOP_K,
|
| 177 |
+
)
|
| 178 |
+
tokens.append(next_token)
|
| 179 |
+
|
| 180 |
+
if next_token == tokenizer.eos_id:
|
| 181 |
+
break
|
| 182 |
+
|
| 183 |
+
if len(tokens) >= CONTEXT:
|
| 184 |
+
print("\n[Context limit reached]")
|
| 185 |
+
break
|
| 186 |
+
|
| 187 |
+
generated_text = tokenizer.decode(tokens)
|
| 188 |
+
|
| 189 |
+
print("\n" + "=" * 60)
|
| 190 |
+
print("Veylon Alpha")
|
| 191 |
+
print("=" * 60)
|
| 192 |
+
print(generated_text)
|
| 193 |
+
print("=" * 60)
|
inspect_weights.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import h5py
|
| 2 |
+
|
| 3 |
+
with h5py.File('generative_model.weights.h5', 'r') as f:
|
| 4 |
+
def print_structure(name, obj):
|
| 5 |
+
if isinstance(obj, h5py.Dataset):
|
| 6 |
+
print(f"{name}: {obj.shape} dtype={obj.dtype}")
|
| 7 |
+
f.visititems(print_structure)
|
model_build.py
ADDED
|
@@ -0,0 +1,56 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# model_build.py
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
os.environ["KERAS_BACKEND"] = "jax"
|
| 5 |
+
|
| 6 |
+
import keras
|
| 7 |
+
|
| 8 |
+
from veylon_model import create_llm
|
| 9 |
+
from tokenizer import TokenizerWrapper
|
| 10 |
+
|
| 11 |
+
from config import (
|
| 12 |
+
CONTEXT,
|
| 13 |
+
ffn_mult,
|
| 14 |
+
d_Latent,
|
| 15 |
+
D_MODEL,
|
| 16 |
+
numberofheads,
|
| 17 |
+
numberoflayers,
|
| 18 |
+
vocab_size
|
| 19 |
+
)
|
| 20 |
+
|
| 21 |
+
# Mixed precision is set inside veylon_model.py, but we ensure it here too
|
| 22 |
+
keras.mixed_precision.set_global_policy("mixed_bfloat16")
|
| 23 |
+
|
| 24 |
+
tokenizer = TokenizerWrapper("tokenizer.json")
|
| 25 |
+
|
| 26 |
+
# n_kv_heads is completely removed - pure MLA
|
| 27 |
+
model = create_llm(
|
| 28 |
+
vocab_size=vocab_size,
|
| 29 |
+
d_model=D_MODEL,
|
| 30 |
+
n_layers=numberoflayers,
|
| 31 |
+
n_heads=numberofheads,
|
| 32 |
+
d_latent=d_Latent,
|
| 33 |
+
ffn_mult=ffn_mult,
|
| 34 |
+
max_seq_len=CONTEXT,
|
| 35 |
+
use_moe=False,
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
optimizer = keras.optimizers.AdamW(
|
| 39 |
+
learning_rate=1e-4,
|
| 40 |
+
weight_decay=0.01,
|
| 41 |
+
global_clipnorm=1.0,
|
| 42 |
+
)
|
| 43 |
+
|
| 44 |
+
loss_fn = keras.losses.SparseCategoricalCrossentropy(
|
| 45 |
+
from_logits=True,
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
model.compile(
|
| 49 |
+
optimizer=optimizer,
|
| 50 |
+
loss=loss_fn,
|
| 51 |
+
jit_compile=True,
|
| 52 |
+
)
|
| 53 |
+
|
| 54 |
+
# Explicitly build to print summary
|
| 55 |
+
model.build((None, CONTEXT))
|
| 56 |
+
model.summary()
|
token_train.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from tokenizer import train_sentencepiece
|
| 3 |
+
from config import vocab_size
|
| 4 |
+
DATA_DIR = "./training_data"
|
| 5 |
+
|
| 6 |
+
text_files = [
|
| 7 |
+
os.path.join(DATA_DIR, f)
|
| 8 |
+
for f in os.listdir(DATA_DIR)
|
| 9 |
+
if f.endswith(".txt")
|
| 10 |
+
]
|
| 11 |
+
|
| 12 |
+
print(f"Found {len(text_files)} text files")
|
| 13 |
+
|
| 14 |
+
if not text_files:
|
| 15 |
+
raise ValueError(
|
| 16 |
+
f"No .txt files found in {DATA_DIR}"
|
| 17 |
+
)
|
| 18 |
+
|
| 19 |
+
tokenizer = train_sentencepiece(text_files,vocab_size=vocab_size)
|
| 20 |
+
|
tokenizer.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fa5baedc3bd88d9f1321fcb36c8d9d73c30727bfafb6276951697df6a20d9028
|
| 3 |
+
size 369778
|
tokenizer.vocab
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
train.py
ADDED
|
@@ -0,0 +1,439 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import time
|
| 6 |
+
|
| 7 |
+
os.environ['KERAS_BACKEND'] = 'jax'
|
| 8 |
+
os.environ['XLA_PYTHON_CLIENT_PREALLOCATE'] = 'true'
|
| 9 |
+
os.environ['XLA_PYTHON_CLIENT_MEM_FRACTION'] = '0.90'
|
| 10 |
+
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'
|
| 11 |
+
os.environ['PYTHONWARNINGS'] = 'ignore'
|
| 12 |
+
os.environ['TF_DATA_EXPERIMENTAL_SLACK'] = '1'
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
import tensorflow as tf
|
| 16 |
+
import jax
|
| 17 |
+
import keras
|
| 18 |
+
|
| 19 |
+
from tokenizer import TokenizerWrapper
|
| 20 |
+
from veylon_model import create_llm
|
| 21 |
+
|
| 22 |
+
try:
|
| 23 |
+
from config import (EPOCHS, CONTEXT, vocab_size, D_MODEL, numberoflayers,
|
| 24 |
+
numberofheads, d_Latent, ffn_mult, swa_window,
|
| 25 |
+
num_kv_heads, learning_rate, weight_decay, batch_size,
|
| 26 |
+
data_path)
|
| 27 |
+
except Exception:
|
| 28 |
+
EPOCHS = 20
|
| 29 |
+
CONTEXT = 2048
|
| 30 |
+
vocab_size = 32000
|
| 31 |
+
D_MODEL = 512
|
| 32 |
+
numberoflayers = 8
|
| 33 |
+
numberofheads = 8
|
| 34 |
+
d_Latent = 128
|
| 35 |
+
ffn_mult = 3.5
|
| 36 |
+
swa_window = 1024
|
| 37 |
+
num_kv_heads = 2
|
| 38 |
+
learning_rate = 1e-4
|
| 39 |
+
weight_decay = 0.01
|
| 40 |
+
batch_size = 8
|
| 41 |
+
data_path = './training_data'
|
| 42 |
+
|
| 43 |
+
keras.mixed_precision.set_global_policy('mixed_bfloat16')
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 47 |
+
# PIPELINE
|
| 48 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 49 |
+
|
| 50 |
+
def precompute_windows(tokens: np.ndarray, seq_len: int, stride: int) -> np.ndarray:
|
| 51 |
+
"""
|
| 52 |
+
Slice the token array into overlapping (seq_len+1) windows using
|
| 53 |
+
sliding_window_view, then stride-subsample.
|
| 54 |
+
|
| 55 |
+
sliding_window_view is preferred over as_strided: it validates bounds
|
| 56 |
+
internally and raises a clear error instead of silently reading garbage
|
| 57 |
+
memory past the array end.
|
| 58 |
+
|
| 59 |
+
The [::stride] slice is a zero-copy view. ascontiguousarray forces one
|
| 60 |
+
real allocation — a dense (N, seq_len+1) int32 buffer — which TensorFlow
|
| 61 |
+
can wrap without copying again.
|
| 62 |
+
|
| 63 |
+
Memory: for 5.7M tokens, CONTEXT=2048, stride=1024 → ~45 MB. Trivial
|
| 64 |
+
against the 42 GB host RAM on Colab/Kaggle TPU instances.
|
| 65 |
+
"""
|
| 66 |
+
window_len = seq_len + 1
|
| 67 |
+
view = np.lib.stride_tricks.sliding_window_view(tokens, window_len)[::stride]
|
| 68 |
+
|
| 69 |
+
if len(view) == 0:
|
| 70 |
+
raise ValueError(
|
| 71 |
+
f"Token array ({len(tokens):,}) is too short for "
|
| 72 |
+
f"seq_len={seq_len}, stride={stride}."
|
| 73 |
+
)
|
| 74 |
+
|
| 75 |
+
return np.ascontiguousarray(view, dtype=np.int32)
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def _make_tf_data_options() -> tf.data.Options:
|
| 79 |
+
"""
|
| 80 |
+
Bundle all tf.data graph-level knobs in one place.
|
| 81 |
+
|
| 82 |
+
map_parallelization: lets the runtime fuse and parallelize map ops across
|
| 83 |
+
the thread pool automatically — no manual num_parallel_calls needed for
|
| 84 |
+
simple element-wise transforms.
|
| 85 |
+
|
| 86 |
+
parallel_batch: batching itself becomes a parallel operation; each batch
|
| 87 |
+
slot is filled concurrently instead of serially.
|
| 88 |
+
|
| 89 |
+
experimental_slack: introduces a one-step slack between the prefetch
|
| 90 |
+
stage and the training loop so the pipeline stays one batch ahead
|
| 91 |
+
without blocking on the next step. Redundant with the env var above
|
| 92 |
+
but belt-and-suspenders costs nothing.
|
| 93 |
+
"""
|
| 94 |
+
opts = tf.data.Options()
|
| 95 |
+
opts.experimental_optimization.map_parallelization = True
|
| 96 |
+
opts.experimental_optimization.parallel_batch = True
|
| 97 |
+
opts.experimental_slack = True
|
| 98 |
+
return opts
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def create_tf_dataset(
|
| 102 |
+
windows: np.ndarray,
|
| 103 |
+
batch_size: int,
|
| 104 |
+
shuffle: bool = True,
|
| 105 |
+
shuffle_seed: int = 42,
|
| 106 |
+
cache_path: str = "", # "" → in-memory cache; "/tmp/..." → disk
|
| 107 |
+
) -> tf.data.Dataset:
|
| 108 |
+
"""
|
| 109 |
+
High-throughput tf.data pipeline.
|
| 110 |
+
|
| 111 |
+
Ordering rationale (each decision is load-bearing):
|
| 112 |
+
|
| 113 |
+
1. from_tensor_slices(windows)
|
| 114 |
+
TensorFlow wraps the numpy buffer directly — no tf.constant() copy.
|
| 115 |
+
Cardinality is exactly N (known), enabling optimal prefetch depth and
|
| 116 |
+
shuffle buffer auto-sizing downstream.
|
| 117 |
+
|
| 118 |
+
2. cache() ← BEFORE shuffle and map
|
| 119 |
+
Caches individual (seq_len+1,) sequence tensors. On epoch 2+ the
|
| 120 |
+
entire dataset is served from RAM, eliminating all numpy I/O.
|
| 121 |
+
Caching before shuffle means the cache stores the canonical data once;
|
| 122 |
+
shuffle order changes every epoch without invalidating the cache.
|
| 123 |
+
Caching after batching would freeze batch composition forever, which
|
| 124 |
+
breaks per-epoch shuffle semantics.
|
| 125 |
+
|
| 126 |
+
3. shuffle() ← AFTER cache, BEFORE repeat
|
| 127 |
+
Operates on individual sequences (correct granularity).
|
| 128 |
+
buffer=min(N, 20_000) bounds peak memory while covering the full
|
| 129 |
+
dataset for corpora this size.
|
| 130 |
+
|
| 131 |
+
4. repeat() ← AFTER shuffle, BEFORE map/batch
|
| 132 |
+
Placed here so the shuffle buffer can draw across epoch boundaries,
|
| 133 |
+
preventing an artificial ordering reset at each epoch edge.
|
| 134 |
+
|
| 135 |
+
5. map(split_xy)
|
| 136 |
+
Slices each (seq_len+1,) window into x=(seq_len,) and y=(seq_len,).
|
| 137 |
+
No @tf.function decorator needed — tf.data traces the lambda itself.
|
| 138 |
+
num_parallel_calls=AUTOTUNE parallelises across the CPU thread pool.
|
| 139 |
+
|
| 140 |
+
6. batch(drop_remainder=True)
|
| 141 |
+
Static shape [batch_size, seq_len] — mandatory for XLA/TPU.
|
| 142 |
+
No symbolic/dynamic shapes ever reach the device.
|
| 143 |
+
|
| 144 |
+
7. prefetch(AUTOTUNE)
|
| 145 |
+
Keeps the device-transfer queue full. Always last so it buffers
|
| 146 |
+
fully-formed batches, not individual sequences.
|
| 147 |
+
|
| 148 |
+
8. with_options(...)
|
| 149 |
+
Applies graph-level optimizations in one shot after the pipeline is
|
| 150 |
+
fully defined.
|
| 151 |
+
"""
|
| 152 |
+
n = len(windows)
|
| 153 |
+
|
| 154 |
+
ds = tf.data.Dataset.from_tensor_slices(windows) # (1) wrap
|
| 155 |
+
|
| 156 |
+
ds = ds.cache(cache_path) # (2) cache
|
| 157 |
+
|
| 158 |
+
if shuffle:
|
| 159 |
+
buf = min(n, 20_000)
|
| 160 |
+
ds = ds.shuffle(buffer_size=buf, seed=shuffle_seed,
|
| 161 |
+
reshuffle_each_iteration=True) # (3) shuffle
|
| 162 |
+
|
| 163 |
+
ds = ds.repeat() # (4) repeat
|
| 164 |
+
|
| 165 |
+
ds = ds.map( # (5) split
|
| 166 |
+
lambda w: (w[:-1], w[1:]),
|
| 167 |
+
num_parallel_calls=tf.data.AUTOTUNE,
|
| 168 |
+
)
|
| 169 |
+
|
| 170 |
+
ds = ds.batch(batch_size, drop_remainder=True) # (6) batch
|
| 171 |
+
|
| 172 |
+
ds = ds.prefetch(tf.data.AUTOTUNE) # (7) prefetch
|
| 173 |
+
|
| 174 |
+
ds = ds.with_options(_make_tf_data_options()) # (8) opts
|
| 175 |
+
|
| 176 |
+
return ds
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 180 |
+
# BENCHMARK
|
| 181 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 182 |
+
|
| 183 |
+
def benchmark_dataset(
|
| 184 |
+
ds: tf.data.Dataset,
|
| 185 |
+
batch_size: int,
|
| 186 |
+
seq_len: int,
|
| 187 |
+
n_batches: int = 200,
|
| 188 |
+
label: str = "dataset",
|
| 189 |
+
) -> dict:
|
| 190 |
+
"""
|
| 191 |
+
Measure pipeline throughput independently of TPU compute.
|
| 192 |
+
|
| 193 |
+
We do NOT call .numpy() inside the loop. That forces a device→host
|
| 194 |
+
transfer and a Python GIL acquisition on every batch, serialising what
|
| 195 |
+
should be a fully-async prefetch chain. Instead we consume the iterator
|
| 196 |
+
and let TensorFlow materialise tensors lazily. The end-to-end wall time
|
| 197 |
+
is what the TPU will actually experience.
|
| 198 |
+
|
| 199 |
+
Interpret results:
|
| 200 |
+
pipeline_tok/s > 2× tpu_tok/s → pipeline is not the bottleneck
|
| 201 |
+
pipeline_tok/s < 1.5× tpu_tok/s → still starving the TPU
|
| 202 |
+
"""
|
| 203 |
+
print(f"\n{'─'*52}")
|
| 204 |
+
print(f"Benchmarking: {label}")
|
| 205 |
+
print(f"{'─'*52}")
|
| 206 |
+
|
| 207 |
+
# Warm up: populate cache + prefetch buffers before timing starts.
|
| 208 |
+
for _ in ds.take(5):
|
| 209 |
+
pass
|
| 210 |
+
print("Warm-up (5 batches) complete.")
|
| 211 |
+
|
| 212 |
+
start = time.perf_counter()
|
| 213 |
+
for _ in ds.take(n_batches):
|
| 214 |
+
pass
|
| 215 |
+
elapsed = time.perf_counter() - start
|
| 216 |
+
|
| 217 |
+
tokens_per_batch = batch_size * seq_len
|
| 218 |
+
total_tokens = n_batches * tokens_per_batch
|
| 219 |
+
batches_sec = n_batches / elapsed
|
| 220 |
+
tokens_sec = total_tokens / elapsed
|
| 221 |
+
|
| 222 |
+
print(f"Batches measured : {n_batches}")
|
| 223 |
+
print(f"Elapsed : {elapsed:.2f} s")
|
| 224 |
+
print(f"Throughput : {batches_sec:.1f} batches/s")
|
| 225 |
+
print(f"Throughput : {tokens_sec:,.0f} tokens/s")
|
| 226 |
+
|
| 227 |
+
tpu_target = 30_000
|
| 228 |
+
headroom = tokens_sec / tpu_target
|
| 229 |
+
if headroom >= 2.0:
|
| 230 |
+
verdict = f"✓ {headroom:.1f}× headroom over TPU — pipeline is NOT the bottleneck."
|
| 231 |
+
elif headroom >= 1.2:
|
| 232 |
+
verdict = f"⚠ {headroom:.1f}× headroom — marginal; increase batch_size or reduce stride."
|
| 233 |
+
else:
|
| 234 |
+
verdict = f"✗ {tokens_sec:,.0f} tok/s < TPU target {tpu_target:,} — still starving the TPU."
|
| 235 |
+
|
| 236 |
+
print(verdict)
|
| 237 |
+
print(f"{'─'*52}\n")
|
| 238 |
+
|
| 239 |
+
return {"batches_sec": batches_sec, "tokens_sec": tokens_sec,
|
| 240 |
+
"headroom_vs_30k": headroom}
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 244 |
+
# THROUGHPUT CALLBACK
|
| 245 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 246 |
+
|
| 247 |
+
class ThroughputCallback(keras.callbacks.Callback):
|
| 248 |
+
"""
|
| 249 |
+
Logs throughput every LOG_INTERVAL steps and at epoch end.
|
| 250 |
+
|
| 251 |
+
Why coarse logging matters on TPU:
|
| 252 |
+
Keras callbacks run on the Python host. JAX dispatches training steps
|
| 253 |
+
asynchronously — Python returns before XLA finishes the kernel. A
|
| 254 |
+
time.perf_counter() call on every step forces a host sync, equivalent
|
| 255 |
+
to inserting jax.block_until_ready() every step and destroying async
|
| 256 |
+
pipelining. At LOG_INTERVAL=50 only ~2% of steps incur this cost.
|
| 257 |
+
"""
|
| 258 |
+
LOG_INTERVAL = 50
|
| 259 |
+
|
| 260 |
+
def on_epoch_begin(self, epoch, logs=None):
|
| 261 |
+
self._epoch_start = time.perf_counter()
|
| 262 |
+
self._interval_start = self._epoch_start
|
| 263 |
+
self._step_count = 0
|
| 264 |
+
|
| 265 |
+
def on_train_batch_end(self, batch, logs=None):
|
| 266 |
+
self._step_count += 1
|
| 267 |
+
if batch > 0 and batch % self.LOG_INTERVAL == 0:
|
| 268 |
+
now = time.perf_counter()
|
| 269 |
+
elapsed = now - self._interval_start
|
| 270 |
+
tps = (self.LOG_INTERVAL * batch_size * CONTEXT) / elapsed
|
| 271 |
+
loss = logs.get('loss', float('nan'))
|
| 272 |
+
print(f" step {batch:5d} | loss {loss:.4f} | {tps:>10,.0f} tok/s")
|
| 273 |
+
self._interval_start = now
|
| 274 |
+
|
| 275 |
+
def on_epoch_end(self, epoch, logs=None):
|
| 276 |
+
elapsed = time.perf_counter() - self._epoch_start
|
| 277 |
+
total_tokens = self._step_count * batch_size * CONTEXT
|
| 278 |
+
val_loss = logs.get('val_loss', float('nan'))
|
| 279 |
+
print(f"\nEpoch {epoch+1} | {elapsed:.1f}s | "
|
| 280 |
+
f"{total_tokens / elapsed:,.0f} tok/s (avg) | "
|
| 281 |
+
f"val_loss={val_loss:.4f}\n")
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 285 |
+
# HELPERS
|
| 286 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 287 |
+
|
| 288 |
+
def load_text_corpus(path: str) -> str:
|
| 289 |
+
p = Path(path)
|
| 290 |
+
files = [p] if p.is_file() else sorted(p.glob('*.txt'))
|
| 291 |
+
if not files:
|
| 292 |
+
raise FileNotFoundError(f'No .txt files found in {path}')
|
| 293 |
+
return ''.join(fp.read_text(encoding='utf-8') for fp in files)
|
| 294 |
+
|
| 295 |
+
|
| 296 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 297 |
+
# MAIN
|
| 298 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 299 |
+
class MemoryCallback(keras.callbacks.Callback):
|
| 300 |
+
def on_epoch_end(self, epoch, logs=None):
|
| 301 |
+
stats = jax.devices()[0].memory_stats()
|
| 302 |
+
|
| 303 |
+
used = stats["bytes_in_use"] / 1024**3
|
| 304 |
+
peak = stats["peak_bytes_in_use"] / 1024**3
|
| 305 |
+
limit = stats["bytes_limit"] / 1024**3
|
| 306 |
+
|
| 307 |
+
print(
|
| 308 |
+
f"\nHBM: {used:.2f}/{limit:.2f} GB "
|
| 309 |
+
f"(Peak: {peak:.2f} GB)"
|
| 310 |
+
)
|
| 311 |
+
def main():
|
| 312 |
+
print('\n' + '=' * 60)
|
| 313 |
+
print('BACKEND & DEVICE VERIFICATION')
|
| 314 |
+
print('=' * 60)
|
| 315 |
+
print(f'Keras backend : {keras.backend.backend()}')
|
| 316 |
+
print(f'JAX devices : {jax.devices()}')
|
| 317 |
+
print('=' * 60 + '\n')
|
| 318 |
+
|
| 319 |
+
# ── Tokenize ──────────────────────────────────────────────────────────
|
| 320 |
+
tokenizer = TokenizerWrapper('tokenizer.model')
|
| 321 |
+
text = load_text_corpus(data_path)
|
| 322 |
+
raw_tokens = np.asarray(
|
| 323 |
+
tokenizer.encode(text, add_bos=True, add_eos=False), dtype=np.int32
|
| 324 |
+
)
|
| 325 |
+
print(f'Total tokens: {len(raw_tokens):,}')
|
| 326 |
+
|
| 327 |
+
split_idx = int(len(raw_tokens) * 0.98)
|
| 328 |
+
train_tokens = raw_tokens[:split_idx]
|
| 329 |
+
val_tokens = raw_tokens[split_idx:]
|
| 330 |
+
|
| 331 |
+
stride = max(1, CONTEXT // 2)
|
| 332 |
+
|
| 333 |
+
# ── Pre-compute windows ────────────────────────────────────────────────
|
| 334 |
+
print("Pre-computing windows...", end=" ", flush=True)
|
| 335 |
+
t0 = time.perf_counter()
|
| 336 |
+
train_windows = precompute_windows(train_tokens, CONTEXT, stride)
|
| 337 |
+
val_windows = precompute_windows(val_tokens, CONTEXT, stride)
|
| 338 |
+
print(f"done in {(time.perf_counter() - t0) * 1000:.1f} ms")
|
| 339 |
+
print(f" train: {train_windows.shape} ({train_windows.nbytes / 1e6:.1f} MB)")
|
| 340 |
+
print(f" val: {val_windows.shape} ({val_windows.nbytes / 1e6:.1f} MB)\n")
|
| 341 |
+
|
| 342 |
+
# ── Build datasets ─────────────────────────────────────────────────────
|
| 343 |
+
train_ds = create_tf_dataset(train_windows, batch_size, shuffle=True)
|
| 344 |
+
val_ds = create_tf_dataset(val_windows, batch_size, shuffle=False)
|
| 345 |
+
|
| 346 |
+
# ── Benchmark before training ──────────────────────────────────────────
|
| 347 |
+
# Verifies the pipeline can outpace the TPU before we waste a run.
|
| 348 |
+
benchmark_dataset(train_ds, batch_size, CONTEXT, n_batches=200,
|
| 349 |
+
label="train_ds")
|
| 350 |
+
|
| 351 |
+
# ── Build model ────────────────────────────────────────────────────────
|
| 352 |
+
os.makedirs('checkpoints', exist_ok=True)
|
| 353 |
+
|
| 354 |
+
model = create_llm(
|
| 355 |
+
vocab_size = tokenizer.vocab_size,
|
| 356 |
+
d_model = D_MODEL,
|
| 357 |
+
n_layers = numberoflayers,
|
| 358 |
+
n_heads = numberofheads,
|
| 359 |
+
d_latent = d_Latent,
|
| 360 |
+
ffn_mult = ffn_mult,
|
| 361 |
+
max_seq_len = CONTEXT,
|
| 362 |
+
use_moe = False,
|
| 363 |
+
num_kv_heads = num_kv_heads,
|
| 364 |
+
swa_window = swa_window,
|
| 365 |
+
)
|
| 366 |
+
|
| 367 |
+
optimizer = keras.optimizers.AdamW(
|
| 368 |
+
learning_rate = learning_rate,
|
| 369 |
+
weight_decay = weight_decay,
|
| 370 |
+
global_clipnorm = 1.0,
|
| 371 |
+
)
|
| 372 |
+
model.compile(
|
| 373 |
+
optimizer = optimizer,
|
| 374 |
+
loss = keras.losses.SparseCategoricalCrossentropy(from_logits=True),
|
| 375 |
+
jit_compile = True,
|
| 376 |
+
)
|
| 377 |
+
|
| 378 |
+
# ── Warmup with exact training shape ──────────────────────────────────
|
| 379 |
+
# XLA traces a kernel per unique input shape. Warming up with any other
|
| 380 |
+
# shape (e.g. the original (1, min(4, CONTEXT))) causes a full recompile
|
| 381 |
+
# on the first real training batch — wasting 30-120 s on the TPU.
|
| 382 |
+
print("Warming up XLA with training shape...", end=" ", flush=True)
|
| 383 |
+
t0 = time.perf_counter()
|
| 384 |
+
dummy_x = np.zeros((batch_size, CONTEXT), dtype=np.int32)
|
| 385 |
+
_ = model(dummy_x, training=False)
|
| 386 |
+
_ = model(dummy_x, training=True)
|
| 387 |
+
print(f"done in {time.perf_counter() - t0:.2f}s\n")
|
| 388 |
+
|
| 389 |
+
# ── Forward pass sanity check ──────────────────────────────────────────
|
| 390 |
+
# One explicit transfer (np.array) then work in NumPy — avoids the 4
|
| 391 |
+
# implicit device→host round trips of the original float(ops.min(...))
|
| 392 |
+
# scalar-cast pattern.
|
| 393 |
+
print('─' * 52)
|
| 394 |
+
print('Forward pass validation')
|
| 395 |
+
print('─' * 52)
|
| 396 |
+
for x_batch, _ in train_ds.take(1):
|
| 397 |
+
logits = model(x_batch[:1], training=False)
|
| 398 |
+
logits_np = np.array(logits)
|
| 399 |
+
print(f'Input : {x_batch.shape}')
|
| 400 |
+
print(f'Logits : {logits.shape}')
|
| 401 |
+
print(f'min={logits_np.min():.4f} max={logits_np.max():.4f}')
|
| 402 |
+
if np.isnan(logits_np).any() or np.isinf(logits_np).any():
|
| 403 |
+
raise ValueError('Forward pass produced NaN/Inf — check model init.')
|
| 404 |
+
print('✓ Forward pass clean.\n')
|
| 405 |
+
|
| 406 |
+
# ── Steps per epoch from actual window count ───────────────────────────
|
| 407 |
+
# Original used len(tokens) // (batch*seq) which undercounts overlapping
|
| 408 |
+
# windows and causes epochs to terminate prematurely.
|
| 409 |
+
print(f"Total params: {model.count_params():,}")
|
| 410 |
+
steps_per_epoch = max(1, len(train_windows) // batch_size)
|
| 411 |
+
validation_steps = max(1, len(val_windows) // batch_size)
|
| 412 |
+
print(f"steps_per_epoch : {steps_per_epoch}")
|
| 413 |
+
print(f"validation_steps : {validation_steps}\n")
|
| 414 |
+
|
| 415 |
+
# ── Train ──────────────────────────────────────────────────────────────
|
| 416 |
+
model.fit(
|
| 417 |
+
train_ds,
|
| 418 |
+
validation_data = val_ds,
|
| 419 |
+
epochs = EPOCHS,
|
| 420 |
+
steps_per_epoch = steps_per_epoch,
|
| 421 |
+
validation_steps = validation_steps,
|
| 422 |
+
callbacks = [
|
| 423 |
+
keras.callbacks.TerminateOnNaN(),
|
| 424 |
+
keras.callbacks.ModelCheckpoint(
|
| 425 |
+
filepath = 'checkpoints/veylon_{epoch:02d}.weights.h5',
|
| 426 |
+
save_freq = 'epoch',
|
| 427 |
+
save_weights_only = True,
|
| 428 |
+
),
|
| 429 |
+
ThroughputCallback(),
|
| 430 |
+
MemoryCallback(),
|
| 431 |
+
],
|
| 432 |
+
)
|
| 433 |
+
|
| 434 |
+
model.save_weights('veylon_final.weights.h5')
|
| 435 |
+
print('Saved veylon_final.weights.h5')
|
| 436 |
+
|
| 437 |
+
|
| 438 |
+
if __name__ == '__main__':
|
| 439 |
+
main()
|
train2.py
ADDED
|
@@ -0,0 +1,547 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import time
|
| 6 |
+
|
| 7 |
+
os.environ['KERAS_BACKEND'] = 'jax'
|
| 8 |
+
os.environ['XLA_PYTHON_CLIENT_PREALLOCATE'] = 'true'
|
| 9 |
+
os.environ['XLA_PYTHON_CLIENT_MEM_FRACTION'] = '0.90'
|
| 10 |
+
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'
|
| 11 |
+
os.environ['PYTHONWARNINGS'] = 'ignore'
|
| 12 |
+
os.environ['TF_DATA_EXPERIMENTAL_SLACK'] = '1'
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
import tensorflow as tf
|
| 16 |
+
import jax
|
| 17 |
+
import jax.sharding as sharding
|
| 18 |
+
from jax.sharding import PartitionSpec as P
|
| 19 |
+
import keras
|
| 20 |
+
|
| 21 |
+
from tokenizer import TokenizerWrapper
|
| 22 |
+
from veylon_model import create_llm
|
| 23 |
+
|
| 24 |
+
try:
|
| 25 |
+
from config import (EPOCHS, CONTEXT, vocab_size, D_MODEL, numberoflayers,
|
| 26 |
+
numberofheads, d_Latent, ffn_mult, swa_window,
|
| 27 |
+
num_kv_heads, learning_rate, weight_decay, batch_size,
|
| 28 |
+
data_path)
|
| 29 |
+
except Exception:
|
| 30 |
+
EPOCHS = 20
|
| 31 |
+
CONTEXT = 2048
|
| 32 |
+
vocab_size = 32000
|
| 33 |
+
D_MODEL = 512
|
| 34 |
+
numberoflayers = 8
|
| 35 |
+
numberofheads = 8
|
| 36 |
+
d_Latent = 128
|
| 37 |
+
ffn_mult = 3.5
|
| 38 |
+
swa_window = 1024
|
| 39 |
+
num_kv_heads = 2
|
| 40 |
+
learning_rate = 1e-4
|
| 41 |
+
weight_decay = 0.01
|
| 42 |
+
batch_size = 8
|
| 43 |
+
data_path = './training_data'
|
| 44 |
+
|
| 45 |
+
keras.mixed_precision.set_global_policy('mixed_bfloat16')
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 49 |
+
# PIPELINE
|
| 50 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 51 |
+
|
| 52 |
+
def precompute_windows(tokens: np.ndarray, seq_len: int, stride: int) -> np.ndarray:
|
| 53 |
+
"""
|
| 54 |
+
Slice the token array into overlapping (seq_len+1) windows using
|
| 55 |
+
sliding_window_view, then stride-subsample.
|
| 56 |
+
|
| 57 |
+
sliding_window_view is preferred over as_strided: it validates bounds
|
| 58 |
+
internally and raises a clear error instead of silently reading garbage
|
| 59 |
+
memory past the array end.
|
| 60 |
+
|
| 61 |
+
The [::stride] slice is a zero-copy view. ascontiguousarray forces one
|
| 62 |
+
real allocation — a dense (N, seq_len+1) int32 buffer — which TensorFlow
|
| 63 |
+
can wrap without copying again.
|
| 64 |
+
|
| 65 |
+
Memory: for 5.7M tokens, CONTEXT=2048, stride=1024 → ~45 MB. Trivial
|
| 66 |
+
against the 42 GB host RAM on Colab/Kaggle TPU instances.
|
| 67 |
+
"""
|
| 68 |
+
window_len = seq_len + 1
|
| 69 |
+
view = np.lib.stride_tricks.sliding_window_view(tokens, window_len)[::stride]
|
| 70 |
+
|
| 71 |
+
if len(view) == 0:
|
| 72 |
+
raise ValueError(
|
| 73 |
+
f"Token array ({len(tokens):,}) is too short for "
|
| 74 |
+
f"seq_len={seq_len}, stride={stride}."
|
| 75 |
+
)
|
| 76 |
+
|
| 77 |
+
return np.ascontiguousarray(view, dtype=np.int32)
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def _make_tf_data_options() -> tf.data.Options:
|
| 81 |
+
"""
|
| 82 |
+
Bundle all tf.data graph-level knobs in one place.
|
| 83 |
+
|
| 84 |
+
map_parallelization: lets the runtime fuse and parallelize map ops across
|
| 85 |
+
the thread pool automatically — no manual num_parallel_calls needed for
|
| 86 |
+
simple element-wise transforms.
|
| 87 |
+
|
| 88 |
+
parallel_batch: batching itself becomes a parallel operation; each batch
|
| 89 |
+
slot is filled concurrently instead of serially.
|
| 90 |
+
|
| 91 |
+
experimental_slack: introduces a one-step slack between the prefetch
|
| 92 |
+
stage and the training loop so the pipeline stays one batch ahead
|
| 93 |
+
without blocking on the next step. Redundant with the env var above
|
| 94 |
+
but belt-and-suspenders costs nothing.
|
| 95 |
+
"""
|
| 96 |
+
opts = tf.data.Options()
|
| 97 |
+
opts.experimental_optimization.map_parallelization = True
|
| 98 |
+
opts.experimental_optimization.parallel_batch = True
|
| 99 |
+
opts.experimental_slack = True
|
| 100 |
+
return opts
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
def create_tf_dataset(
|
| 104 |
+
windows: np.ndarray,
|
| 105 |
+
batch_size: int,
|
| 106 |
+
shuffle: bool = True,
|
| 107 |
+
shuffle_seed: int = 42,
|
| 108 |
+
cache_path: str = "", # "" → in-memory cache; "/tmp/..." → disk
|
| 109 |
+
) -> tf.data.Dataset:
|
| 110 |
+
"""
|
| 111 |
+
High-throughput tf.data pipeline.
|
| 112 |
+
|
| 113 |
+
Ordering rationale (each decision is load-bearing):
|
| 114 |
+
|
| 115 |
+
1. from_tensor_slices(windows)
|
| 116 |
+
TensorFlow wraps the numpy buffer directly — no tf.constant() copy.
|
| 117 |
+
Cardinality is exactly N (known), enabling optimal prefetch depth and
|
| 118 |
+
shuffle buffer auto-sizing downstream.
|
| 119 |
+
|
| 120 |
+
2. cache() ← BEFORE shuffle and map
|
| 121 |
+
Caches individual (seq_len+1,) sequence tensors. On epoch 2+ the
|
| 122 |
+
entire dataset is served from RAM, eliminating all numpy I/O.
|
| 123 |
+
Caching before shuffle means the cache stores the canonical data once;
|
| 124 |
+
shuffle order changes every epoch without invalidating the cache.
|
| 125 |
+
Caching after batching would freeze batch composition forever, which
|
| 126 |
+
breaks per-epoch shuffle semantics.
|
| 127 |
+
|
| 128 |
+
3. shuffle() ← AFTER cache, BEFORE repeat
|
| 129 |
+
Operates on individual sequences (correct granularity).
|
| 130 |
+
buffer=min(N, 20_000) bounds peak memory while covering the full
|
| 131 |
+
dataset for corpora this size.
|
| 132 |
+
|
| 133 |
+
4. repeat() ← AFTER shuffle, BEFORE map/batch
|
| 134 |
+
Placed here so the shuffle buffer can draw across epoch boundaries,
|
| 135 |
+
preventing an artificial ordering reset at each epoch edge.
|
| 136 |
+
|
| 137 |
+
5. map(split_xy)
|
| 138 |
+
Slices each (seq_len+1,) window into x=(seq_len,) and y=(seq_len,).
|
| 139 |
+
No @tf.function decorator needed — tf.data traces the lambda itself.
|
| 140 |
+
num_parallel_calls=AUTOTUNE parallelises across the CPU thread pool.
|
| 141 |
+
|
| 142 |
+
6. batch(drop_remainder=True)
|
| 143 |
+
Static shape [batch_size, seq_len] — mandatory for XLA/TPU.
|
| 144 |
+
No symbolic/dynamic shapes ever reach the device.
|
| 145 |
+
|
| 146 |
+
7. prefetch(AUTOTUNE)
|
| 147 |
+
Keeps the device-transfer queue full. Always last so it buffers
|
| 148 |
+
fully-formed batches, not individual sequences.
|
| 149 |
+
|
| 150 |
+
8. with_options(...)
|
| 151 |
+
Applies graph-level optimizations in one shot after the pipeline is
|
| 152 |
+
fully defined.
|
| 153 |
+
"""
|
| 154 |
+
n = len(windows)
|
| 155 |
+
|
| 156 |
+
ds = tf.data.Dataset.from_tensor_slices(windows) # (1) wrap
|
| 157 |
+
|
| 158 |
+
ds = ds.cache(cache_path) # (2) cache
|
| 159 |
+
|
| 160 |
+
if shuffle:
|
| 161 |
+
buf = min(n, 20_000)
|
| 162 |
+
ds = ds.shuffle(buffer_size=buf, seed=shuffle_seed,
|
| 163 |
+
reshuffle_each_iteration=True) # (3) shuffle
|
| 164 |
+
|
| 165 |
+
ds = ds.repeat() # (4) repeat
|
| 166 |
+
|
| 167 |
+
ds = ds.map( # (5) split
|
| 168 |
+
lambda w: (w[:-1], w[1:]),
|
| 169 |
+
num_parallel_calls=tf.data.AUTOTUNE,
|
| 170 |
+
)
|
| 171 |
+
|
| 172 |
+
ds = ds.batch(batch_size, drop_remainder=True) # (6) batch
|
| 173 |
+
|
| 174 |
+
ds = ds.prefetch(tf.data.AUTOTUNE) # (7) prefetch
|
| 175 |
+
|
| 176 |
+
ds = ds.with_options(_make_tf_data_options()) # (8) opts
|
| 177 |
+
|
| 178 |
+
return ds
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 182 |
+
# BENCHMARK
|
| 183 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 184 |
+
|
| 185 |
+
def benchmark_dataset(
|
| 186 |
+
ds: tf.data.Dataset,
|
| 187 |
+
batch_size: int,
|
| 188 |
+
seq_len: int,
|
| 189 |
+
n_batches: int = 200,
|
| 190 |
+
label: str = "dataset",
|
| 191 |
+
) -> dict:
|
| 192 |
+
"""
|
| 193 |
+
Measure pipeline throughput independently of TPU compute.
|
| 194 |
+
|
| 195 |
+
We do NOT call .numpy() inside the loop. That forces a device→host
|
| 196 |
+
transfer and a Python GIL acquisition on every batch, serialising what
|
| 197 |
+
should be a fully-async prefetch chain. Instead we consume the iterator
|
| 198 |
+
and let TensorFlow materialise tensors lazily. The end-to-end wall time
|
| 199 |
+
is what the TPU will actually experience.
|
| 200 |
+
|
| 201 |
+
Interpret results:
|
| 202 |
+
pipeline_tok/s > 2× tpu_tok/s → pipeline is not the bottleneck
|
| 203 |
+
pipeline_tok/s < 1.5× tpu_tok/s → still starving the TPU
|
| 204 |
+
"""
|
| 205 |
+
print(f"\n{'─'*52}")
|
| 206 |
+
print(f"Benchmarking: {label}")
|
| 207 |
+
print(f"{'─'*52}")
|
| 208 |
+
|
| 209 |
+
# Warm up: populate cache + prefetch buffers before timing starts.
|
| 210 |
+
for _ in ds.take(5):
|
| 211 |
+
pass
|
| 212 |
+
print("Warm-up (5 batches) complete.")
|
| 213 |
+
|
| 214 |
+
start = time.perf_counter()
|
| 215 |
+
for _ in ds.take(n_batches):
|
| 216 |
+
pass
|
| 217 |
+
elapsed = time.perf_counter() - start
|
| 218 |
+
|
| 219 |
+
tokens_per_batch = batch_size * seq_len
|
| 220 |
+
total_tokens = n_batches * tokens_per_batch
|
| 221 |
+
batches_sec = n_batches / elapsed
|
| 222 |
+
tokens_sec = total_tokens / elapsed
|
| 223 |
+
|
| 224 |
+
print(f"Batches measured : {n_batches}")
|
| 225 |
+
print(f"Elapsed : {elapsed:.2f} s")
|
| 226 |
+
print(f"Throughput : {batches_sec:.1f} batches/s")
|
| 227 |
+
print(f"Throughput : {tokens_sec:,.0f} tokens/s")
|
| 228 |
+
|
| 229 |
+
tpu_target = 30_000
|
| 230 |
+
headroom = tokens_sec / tpu_target
|
| 231 |
+
if headroom >= 2.0:
|
| 232 |
+
verdict = f"✓ {headroom:.1f}× headroom over TPU — pipeline is NOT the bottleneck."
|
| 233 |
+
elif headroom >= 1.2:
|
| 234 |
+
verdict = f"⚠ {headroom:.1f}× headroom — marginal; increase batch_size or reduce stride."
|
| 235 |
+
else:
|
| 236 |
+
verdict = f"✗ {tokens_sec:,.0f} tok/s < TPU target {tpu_target:,} — still starving the TPU."
|
| 237 |
+
|
| 238 |
+
print(verdict)
|
| 239 |
+
print(f"{'─'*52}\n")
|
| 240 |
+
|
| 241 |
+
return {"batches_sec": batches_sec, "tokens_sec": tokens_sec,
|
| 242 |
+
"headroom_vs_30k": headroom}
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 246 |
+
# THROUGHPUT CALLBACK
|
| 247 |
+
# ──────────────────────────���──────────────────────────────────────────────────
|
| 248 |
+
|
| 249 |
+
class ThroughputCallback(keras.callbacks.Callback):
|
| 250 |
+
"""
|
| 251 |
+
Logs throughput every LOG_INTERVAL steps and at epoch end.
|
| 252 |
+
|
| 253 |
+
Why coarse logging matters on TPU:
|
| 254 |
+
Keras callbacks run on the Python host. JAX dispatches training steps
|
| 255 |
+
asynchronously — Python returns before XLA finishes the kernel. A
|
| 256 |
+
time.perf_counter() call on every step forces a host sync, equivalent
|
| 257 |
+
to inserting jax.block_until_ready() every step and destroying async
|
| 258 |
+
pipelining. At LOG_INTERVAL=50 only ~2% of steps incur this cost.
|
| 259 |
+
"""
|
| 260 |
+
LOG_INTERVAL = 50
|
| 261 |
+
|
| 262 |
+
def on_epoch_begin(self, epoch, logs=None):
|
| 263 |
+
self._epoch_start = time.perf_counter()
|
| 264 |
+
self._interval_start = self._epoch_start
|
| 265 |
+
self._step_count = 0
|
| 266 |
+
|
| 267 |
+
def on_train_batch_end(self, batch, logs=None):
|
| 268 |
+
self._step_count += 1
|
| 269 |
+
if batch > 0 and batch % self.LOG_INTERVAL == 0:
|
| 270 |
+
now = time.perf_counter()
|
| 271 |
+
elapsed = now - self._interval_start
|
| 272 |
+
# Multiply by effective_batch if data parallelism is active
|
| 273 |
+
tps = (self.LOG_INTERVAL * effective_batch * CONTEXT) / elapsed
|
| 274 |
+
loss = logs.get('loss', float('nan'))
|
| 275 |
+
print(f" step {batch:5d} | loss {loss:.4f} | {tps:>10,.0f} tok/s")
|
| 276 |
+
self._interval_start = now
|
| 277 |
+
|
| 278 |
+
def on_epoch_end(self, epoch, logs=None):
|
| 279 |
+
elapsed = time.perf_counter() - self._epoch_start
|
| 280 |
+
total_tokens = self._step_count * effective_batch * CONTEXT
|
| 281 |
+
val_loss = logs.get('val_loss', float('nan'))
|
| 282 |
+
print(f"\nEpoch {epoch+1} | {elapsed:.1f}s | "
|
| 283 |
+
f"{total_tokens / elapsed:,.0f} tok/s (avg) | "
|
| 284 |
+
f"val_loss={val_loss:.4f}\n")
|
| 285 |
+
|
| 286 |
+
|
| 287 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 288 |
+
# HELPERS
|
| 289 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 290 |
+
|
| 291 |
+
def load_text_corpus(path: str) -> str:
|
| 292 |
+
p = Path(path)
|
| 293 |
+
files = [p] if p.is_file() else sorted(p.glob('*.txt'))
|
| 294 |
+
if not files:
|
| 295 |
+
raise FileNotFoundError(f'No .txt files found in {path}')
|
| 296 |
+
return ''.join(fp.read_text(encoding='utf-8') for fp in files)
|
| 297 |
+
|
| 298 |
+
|
| 299 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 300 |
+
# SHARDING (Data Parallelism for v5e-8)
|
| 301 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 302 |
+
|
| 303 |
+
def setup_data_parallelism():
|
| 304 |
+
"""
|
| 305 |
+
Configure JAX for pure data parallelism across all TPU devices.
|
| 306 |
+
|
| 307 |
+
For v5e-8 (8 chips):
|
| 308 |
+
- Each chip runs the full model
|
| 309 |
+
- Different batch shards go to each chip
|
| 310 |
+
- Gradients are AllReduced across chips
|
| 311 |
+
- Effective batch = batch_per_chip × num_chips
|
| 312 |
+
|
| 313 |
+
This is the simplest and most efficient sharding strategy for models
|
| 314 |
+
that fit on a single chip. For 8M parameters, the model easily fits
|
| 315 |
+
on one v5e chip (16 GB HBM each).
|
| 316 |
+
|
| 317 |
+
Returns mesh_shape (tuple) and sharding spec dict for inputs/weights.
|
| 318 |
+
"""
|
| 319 |
+
devices = jax.devices()
|
| 320 |
+
n_devices = len(devices)
|
| 321 |
+
|
| 322 |
+
if n_devices == 1:
|
| 323 |
+
print(f"Single device detected ({devices[0]}). Data parallelism disabled.")
|
| 324 |
+
return None, None
|
| 325 |
+
|
| 326 |
+
print(f"\n{'─'*52}")
|
| 327 |
+
print(f"Data Parallelism Setup (v5e-{n_devices})")
|
| 328 |
+
print(f"{'─'*52}")
|
| 329 |
+
print(f"Devices: {n_devices}")
|
| 330 |
+
print(f"Device type: {devices[0].platform}")
|
| 331 |
+
|
| 332 |
+
# Create 1D mesh for data parallelism: each axis is one device
|
| 333 |
+
mesh = sharding.Mesh(
|
| 334 |
+
devices=np.array(devices).reshape((n_devices,)),
|
| 335 |
+
axis_names=("batch",)
|
| 336 |
+
)
|
| 337 |
+
|
| 338 |
+
# Sharding specs for inputs and weights
|
| 339 |
+
# Input shape: (batch, seq_len)
|
| 340 |
+
# Shard batch dimension across devices, keep seq_len replicated
|
| 341 |
+
input_spec = P("batch", None)
|
| 342 |
+
|
| 343 |
+
# Weights: replicate across all devices (each chip has full model)
|
| 344 |
+
weight_spec = P(None)
|
| 345 |
+
|
| 346 |
+
# Activation gradients during backward pass
|
| 347 |
+
# Same as input: shard batch, replicate seq
|
| 348 |
+
activation_spec = P("batch", None)
|
| 349 |
+
|
| 350 |
+
print(f"Mesh shape: {mesh.shape}")
|
| 351 |
+
print(f"Input sharding: {input_spec} (batch sharded, seq replicated)")
|
| 352 |
+
print(f"Weight sharding: {weight_spec} (fully replicated)")
|
| 353 |
+
print(f"Expected effective batch: {n_devices} × batch_size")
|
| 354 |
+
print(f"{'─'*52}\n")
|
| 355 |
+
|
| 356 |
+
return mesh, {
|
| 357 |
+
"input": input_spec,
|
| 358 |
+
"weight": weight_spec,
|
| 359 |
+
"activation": activation_spec,
|
| 360 |
+
}
|
| 361 |
+
|
| 362 |
+
|
| 363 |
+
def apply_sharding_to_model(model, mesh, sharding_specs):
|
| 364 |
+
"""
|
| 365 |
+
Apply JAX sharding annotations to a Keras model compiled with JAX backend.
|
| 366 |
+
|
| 367 |
+
WARNING: This is JAX-level sharding and only works if:
|
| 368 |
+
1. Model is built with JAX-native ops (Keras layers with JAX backend)
|
| 369 |
+
2. jit_compile=True is set in model.compile()
|
| 370 |
+
3. No TensorFlow-only ops are used in the model
|
| 371 |
+
|
| 372 |
+
For Keras models, we set the sharding via jax.Array.with_sharding_constraint()
|
| 373 |
+
in a custom train step. However, Keras 3 makes this tricky because it wraps
|
| 374 |
+
training in its own jit.
|
| 375 |
+
|
| 376 |
+
Simpler approach: set jax.config to globally use this mesh, and let XLA
|
| 377 |
+
infer sharding from the mesh context.
|
| 378 |
+
"""
|
| 379 |
+
if mesh is None:
|
| 380 |
+
return
|
| 381 |
+
|
| 382 |
+
# Tell JAX to use this mesh for all operations
|
| 383 |
+
with mesh:
|
| 384 |
+
print("Sharding configuration loaded.")
|
| 385 |
+
print(f" All matmul ops will shard batch dim across {mesh.shape[0]} devices")
|
| 386 |
+
print(f" Gradient AllReduce will use {mesh.shape[0]}-way ICI collective\n")
|
| 387 |
+
|
| 388 |
+
|
| 389 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 390 |
+
# MAIN
|
| 391 |
+
# ─────────────────────────────────────────────────────────────────────────────
|
| 392 |
+
class MemoryCallback(keras.callbacks.Callback):
|
| 393 |
+
def on_epoch_end(self, epoch, logs=None):
|
| 394 |
+
stats = jax.devices()[0].memory_stats()
|
| 395 |
+
|
| 396 |
+
used = stats["bytes_in_use"] / 1024**3
|
| 397 |
+
peak = stats["peak_bytes_in_use"] / 1024**3
|
| 398 |
+
limit = stats["bytes_limit"] / 1024**3
|
| 399 |
+
|
| 400 |
+
print(
|
| 401 |
+
f"\nHBM: {used:.2f}/{limit:.2f} GB "
|
| 402 |
+
f"(Peak: {peak:.2f} GB)"
|
| 403 |
+
)
|
| 404 |
+
def main():
|
| 405 |
+
print('\n' + '=' * 60)
|
| 406 |
+
print('BACKEND & DEVICE VERIFICATION')
|
| 407 |
+
print('=' * 60)
|
| 408 |
+
print(f'Keras backend : {keras.backend.backend()}')
|
| 409 |
+
print(f'JAX devices : {jax.devices()}')
|
| 410 |
+
print('=' * 60 + '\n')
|
| 411 |
+
|
| 412 |
+
# ── Setup data parallelism ─────────────────────────────────────────────
|
| 413 |
+
mesh, sharding_specs = setup_data_parallelism()
|
| 414 |
+
if mesh is not None:
|
| 415 |
+
apply_sharding_to_model(None, mesh, sharding_specs)
|
| 416 |
+
|
| 417 |
+
# ── Tokenize ──────────────────────────────────────────────────────────
|
| 418 |
+
tokenizer = TokenizerWrapper('tokenizer.model')
|
| 419 |
+
text = load_text_corpus(data_path)
|
| 420 |
+
raw_tokens = np.asarray(
|
| 421 |
+
tokenizer.encode(text, add_bos=True, add_eos=False), dtype=np.int32
|
| 422 |
+
)
|
| 423 |
+
print(f'Total tokens: {len(raw_tokens):,}')
|
| 424 |
+
|
| 425 |
+
split_idx = int(len(raw_tokens) * 0.98)
|
| 426 |
+
train_tokens = raw_tokens[:split_idx]
|
| 427 |
+
val_tokens = raw_tokens[split_idx:]
|
| 428 |
+
|
| 429 |
+
stride = max(1, CONTEXT // 2)
|
| 430 |
+
|
| 431 |
+
# ── Pre-compute windows ────────────────────────────────────────────────
|
| 432 |
+
print("Pre-computing windows...", end=" ", flush=True)
|
| 433 |
+
t0 = time.perf_counter()
|
| 434 |
+
train_windows = precompute_windows(train_tokens, CONTEXT, stride)
|
| 435 |
+
val_windows = precompute_windows(val_tokens, CONTEXT, stride)
|
| 436 |
+
print(f"done in {(time.perf_counter() - t0) * 1000:.1f} ms")
|
| 437 |
+
print(f" train: {train_windows.shape} ({train_windows.nbytes / 1e6:.1f} MB)")
|
| 438 |
+
print(f" val: {val_windows.shape} ({val_windows.nbytes / 1e6:.1f} MB)\n")
|
| 439 |
+
|
| 440 |
+
# ── Build datasets ─────────────────────────────────────────────────────
|
| 441 |
+
# If data parallelism is active, the effective batch becomes:
|
| 442 |
+
# batch_per_device × num_devices
|
| 443 |
+
# But tf.data still sees batch_per_device — JAX handles the replication.
|
| 444 |
+
num_devices = len(jax.devices())
|
| 445 |
+
actual_batch_size = batch_size # Per-device batch
|
| 446 |
+
effective_batch = actual_batch_size * num_devices if mesh is not None else actual_batch_size
|
| 447 |
+
|
| 448 |
+
if mesh is not None:
|
| 449 |
+
print(f"Effective batch size: {actual_batch_size} × {num_devices} devices = {effective_batch}")
|
| 450 |
+
|
| 451 |
+
train_ds = create_tf_dataset(train_windows, actual_batch_size, shuffle=True)
|
| 452 |
+
val_ds = create_tf_dataset(val_windows, actual_batch_size, shuffle=False)
|
| 453 |
+
|
| 454 |
+
# ── Benchmark before training ────────────────────────────────────────��─
|
| 455 |
+
# Verifies the pipeline can outpace the TPU before we waste a run.
|
| 456 |
+
benchmark_dataset(train_ds, batch_size, CONTEXT, n_batches=200,
|
| 457 |
+
label="train_ds")
|
| 458 |
+
|
| 459 |
+
# ── Build model ────────────────────────────────────────────────────────
|
| 460 |
+
os.makedirs('checkpoints', exist_ok=True)
|
| 461 |
+
|
| 462 |
+
model = create_llm(
|
| 463 |
+
vocab_size = tokenizer.vocab_size,
|
| 464 |
+
d_model = D_MODEL,
|
| 465 |
+
n_layers = numberoflayers,
|
| 466 |
+
n_heads = numberofheads,
|
| 467 |
+
d_latent = d_Latent,
|
| 468 |
+
ffn_mult = ffn_mult,
|
| 469 |
+
max_seq_len = CONTEXT,
|
| 470 |
+
use_moe = False,
|
| 471 |
+
num_kv_heads = num_kv_heads,
|
| 472 |
+
swa_window = swa_window,
|
| 473 |
+
)
|
| 474 |
+
|
| 475 |
+
optimizer = keras.optimizers.AdamW(
|
| 476 |
+
learning_rate = learning_rate,
|
| 477 |
+
weight_decay = weight_decay,
|
| 478 |
+
global_clipnorm = 1.0,
|
| 479 |
+
)
|
| 480 |
+
model.compile(
|
| 481 |
+
optimizer = optimizer,
|
| 482 |
+
loss = keras.losses.SparseCategoricalCrossentropy(from_logits=True),
|
| 483 |
+
jit_compile = True,
|
| 484 |
+
)
|
| 485 |
+
|
| 486 |
+
# ── Warmup with exact training shape ──────────────────────────────────
|
| 487 |
+
# XLA traces a kernel per unique input shape. Warming up with any other
|
| 488 |
+
# shape (e.g. the original (1, min(4, CONTEXT))) causes a full recompile
|
| 489 |
+
# on the first real training batch — wasting 30-120 s on the TPU.
|
| 490 |
+
print("Warming up XLA with training shape...", end=" ", flush=True)
|
| 491 |
+
t0 = time.perf_counter()
|
| 492 |
+
dummy_x = np.zeros((batch_size, CONTEXT), dtype=np.int32)
|
| 493 |
+
_ = model(dummy_x, training=False)
|
| 494 |
+
_ = model(dummy_x, training=True)
|
| 495 |
+
print(f"done in {time.perf_counter() - t0:.2f}s\n")
|
| 496 |
+
|
| 497 |
+
# ── Forward pass sanity check ──────────────────────────────────────────
|
| 498 |
+
# One explicit transfer (np.array) then work in NumPy — avoids the 4
|
| 499 |
+
# implicit device→host round trips of the original float(ops.min(...))
|
| 500 |
+
# scalar-cast pattern.
|
| 501 |
+
print('─' * 52)
|
| 502 |
+
print('Forward pass validation')
|
| 503 |
+
print('─' * 52)
|
| 504 |
+
for x_batch, _ in train_ds.take(1):
|
| 505 |
+
logits = model(x_batch[:1], training=False)
|
| 506 |
+
logits_np = np.array(logits)
|
| 507 |
+
print(f'Input : {x_batch.shape}')
|
| 508 |
+
print(f'Logits : {logits.shape}')
|
| 509 |
+
print(f'min={logits_np.min():.4f} max={logits_np.max():.4f}')
|
| 510 |
+
if np.isnan(logits_np).any() or np.isinf(logits_np).any():
|
| 511 |
+
raise ValueError('Forward pass produced NaN/Inf — check model init.')
|
| 512 |
+
print('✓ Forward pass clean.\n')
|
| 513 |
+
|
| 514 |
+
# ── Steps per epoch from actual window count ───────────────────────────
|
| 515 |
+
# Original used len(tokens) // (batch*seq) which undercounts overlapping
|
| 516 |
+
# windows and causes epochs to terminate prematurely.
|
| 517 |
+
print(f"Total params: {model.count_params():,}")
|
| 518 |
+
steps_per_epoch = max(1, len(train_windows) // batch_size)
|
| 519 |
+
validation_steps = max(1, len(val_windows) // batch_size)
|
| 520 |
+
print(f"steps_per_epoch : {steps_per_epoch}")
|
| 521 |
+
print(f"validation_steps : {validation_steps}\n")
|
| 522 |
+
|
| 523 |
+
# ── Train ──────────────────────────────────────────────────────────────
|
| 524 |
+
model.fit(
|
| 525 |
+
train_ds,
|
| 526 |
+
validation_data = val_ds,
|
| 527 |
+
epochs = EPOCHS,
|
| 528 |
+
steps_per_epoch = steps_per_epoch,
|
| 529 |
+
validation_steps = validation_steps,
|
| 530 |
+
callbacks = [
|
| 531 |
+
keras.callbacks.TerminateOnNaN(),
|
| 532 |
+
keras.callbacks.ModelCheckpoint(
|
| 533 |
+
filepath = 'checkpoints/veylon_{epoch:02d}.weights.h5',
|
| 534 |
+
save_freq = 'epoch',
|
| 535 |
+
save_weights_only = True,
|
| 536 |
+
),
|
| 537 |
+
ThroughputCallback(),
|
| 538 |
+
MemoryCallback(),
|
| 539 |
+
],
|
| 540 |
+
)
|
| 541 |
+
|
| 542 |
+
model.save_weights('veylon_final.weights.h5')
|
| 543 |
+
print('Saved veylon_final.weights.h5')
|
| 544 |
+
|
| 545 |
+
|
| 546 |
+
if __name__ == '__main__':
|
| 547 |
+
main()
|
veylon_attention.py
ADDED
|
@@ -0,0 +1,640 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
veylon_attention.py — Veylon Alpha 1
|
| 3 |
+
======================================
|
| 4 |
+
Block-tiled Native GQA Sliding Window Attention for JAX / Keras-3 / TPU v5e.
|
| 5 |
+
|
| 6 |
+
Why NOT lax.scan + dynamic_slice
|
| 7 |
+
---------------------------------
|
| 8 |
+
The previous implementation used:
|
| 9 |
+
|
| 10 |
+
jax.lax.scan(step, None, (jnp.arange(S), q_scan))
|
| 11 |
+
|
| 12 |
+
with `dynamic_slice(k_pad, (0,0,t,0), (B,Hkv,W,D))` inside the body.
|
| 13 |
+
|
| 14 |
+
XLA/XLA-TPU has a known pathology with this pattern:
|
| 15 |
+
- The scan body contains a dynamic gather (dynamic_slice on a traced int `t`).
|
| 16 |
+
- XLA's while-loop lowering stages ALL window materializations into a single
|
| 17 |
+
large buffer to enable pipeline prefetching.
|
| 18 |
+
- On TPU this produces a hidden tensor [S, B, Hkv, W, D] which is exactly
|
| 19 |
+
what you see in the OOM allocation log: f32[1024,8,4,512,64].
|
| 20 |
+
- jax.checkpoint does NOT protect against this because it is a compiler
|
| 21 |
+
(HLO-level) allocation, not a JAX-level rematerialization artifact.
|
| 22 |
+
|
| 23 |
+
The fix: block-tiled computation with fully static tensor shapes
|
| 24 |
+
----------------------------------------------------------------
|
| 25 |
+
Instead of iterating over S individual tokens we iterate over
|
| 26 |
+
(S / BLK) blocks of queries. For each block:
|
| 27 |
+
|
| 28 |
+
- q_blk : [B, Hq, BLK, D] <- static shape
|
| 29 |
+
- k_blk : [B, Hkv, BLK + W - 1, D] <- static shape
|
| 30 |
+
- scores : [B, Hkv, G, BLK, BLK+W-1] <- static shape
|
| 31 |
+
|
| 32 |
+
XLA sees ONLY static shapes inside the map body. There is no
|
| 33 |
+
gather-inside-loop pattern. XLA can freely fuse, pipeline, and
|
| 34 |
+
tile these einsums onto TPU systolic arrays without hidden buffers.
|
| 35 |
+
|
| 36 |
+
lax.map vs lax.scan
|
| 37 |
+
--------------------
|
| 38 |
+
We use `jax.lax.map` (not `lax.scan`) because:
|
| 39 |
+
- Each block is fully independent — no carry state is needed.
|
| 40 |
+
- `lax.map` lowers to a while_loop with no accumulation buffer.
|
| 41 |
+
- `lax.scan` always allocates an output stacked along the scan axis;
|
| 42 |
+
lax.map lets XLA write directly into the preallocated output slice.
|
| 43 |
+
- This eliminates the [n_blocks, B, Hq, BLK, D] intermediate stack.
|
| 44 |
+
(We still do one reshape at the end, but that is a view, not a copy.)
|
| 45 |
+
|
| 46 |
+
Tensor size invariants
|
| 47 |
+
----------------------
|
| 48 |
+
NEVER created:
|
| 49 |
+
[B, Hq, S, S] full attention score matrix
|
| 50 |
+
[S, B, Hkv, W, D] per-token window stack (old scan bug)
|
| 51 |
+
[B, Hq, Hkv, S, D] duplicated KV heads
|
| 52 |
+
|
| 53 |
+
Largest tensors inside map body (all STATIC shapes):
|
| 54 |
+
k_blk, v_blk : [B, Hkv, BLK+W-1, D]
|
| 55 |
+
scores : [B, Hkv, G, BLK, BLK+W-1]
|
| 56 |
+
probs : same
|
| 57 |
+
|
| 58 |
+
Memory complexity
|
| 59 |
+
-----------------
|
| 60 |
+
k_pad, v_pad : O(B × Hkv × (S + W) × D) linear in S, scales with Hkv
|
| 61 |
+
q : O(B × Hq × S × D)
|
| 62 |
+
Per block : O(B × Hkv × (BLK + W) × D) constant w.r.t. S
|
| 63 |
+
Output : O(B × Hq × S × D)
|
| 64 |
+
TOTAL : O(B × (Hkv + Hq) × S × D) strictly linear in S
|
| 65 |
+
|
| 66 |
+
TPU-specific notes
|
| 67 |
+
------------------
|
| 68 |
+
* All shapes in map body are concrete at trace time — XLA never inserts
|
| 69 |
+
shape-dependent conditionals or recompiles.
|
| 70 |
+
* BF16 inputs are cast to FP32 before matmul/softmax, then cast back.
|
| 71 |
+
* BLK should be a multiple of 128 on TPU v5e for optimal systolic array
|
| 72 |
+
utilization (default 128, tunable via `block_size` parameter).
|
| 73 |
+
* The precomputed mask is a static bool array passed as a closed-over
|
| 74 |
+
constant — XLA fuses it into the einsum kernel at no extra memory cost.
|
| 75 |
+
* jax.checkpoint on the map body protects backward-pass activations
|
| 76 |
+
(one block at a time, not one token at a time — much coarser remat).
|
| 77 |
+
"""
|
| 78 |
+
|
| 79 |
+
from __future__ import annotations
|
| 80 |
+
|
| 81 |
+
import math
|
| 82 |
+
from functools import partial
|
| 83 |
+
from typing import Optional
|
| 84 |
+
|
| 85 |
+
import jax
|
| 86 |
+
import jax.numpy as jnp
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
# ---------------------------------------------------------------------------
|
| 90 |
+
# Public tuning constant
|
| 91 |
+
# ---------------------------------------------------------------------------
|
| 92 |
+
|
| 93 |
+
# Default block size for TPU v5e. Must be a power of 2 and ≥ 1.
|
| 94 |
+
# Larger blocks → fewer kernel launches, better systolic utilisation.
|
| 95 |
+
# Smaller blocks → lower peak HBM per block (useful when W is huge).
|
| 96 |
+
TPU_BLOCK_SIZE: int = 128
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
# ---------------------------------------------------------------------------
|
| 100 |
+
# Utility helpers
|
| 101 |
+
# ---------------------------------------------------------------------------
|
| 102 |
+
|
| 103 |
+
def _next_power_of_2(x: int) -> int:
|
| 104 |
+
x = int(x)
|
| 105 |
+
if x <= 1:
|
| 106 |
+
return 1
|
| 107 |
+
return 1 << (x - 1).bit_length()
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
def apply_rope(
|
| 111 |
+
x: jnp.ndarray,
|
| 112 |
+
cos: jnp.ndarray,
|
| 113 |
+
sin: jnp.ndarray,
|
| 114 |
+
offset: int = 0,
|
| 115 |
+
) -> jnp.ndarray:
|
| 116 |
+
"""
|
| 117 |
+
Apply Rotary Position Encoding to [..., S, D].
|
| 118 |
+
cos/sin tables are pre-built for D//2 (half the head dim).
|
| 119 |
+
Uses dynamic_slice — safe under jit and XLA outside of attention body.
|
| 120 |
+
"""
|
| 121 |
+
if x.ndim < 2:
|
| 122 |
+
raise ValueError(f"apply_rope: expected ≥2-D input, got shape {x.shape}")
|
| 123 |
+
d = int(x.shape[-1])
|
| 124 |
+
if d % 2 != 0:
|
| 125 |
+
raise ValueError(f"apply_rope: head dim must be even, got {d}")
|
| 126 |
+
half = d // 2
|
| 127 |
+
seq_len = int(x.shape[-2])
|
| 128 |
+
if offset + seq_len > int(cos.shape[0]):
|
| 129 |
+
raise ValueError(
|
| 130 |
+
f"apply_rope: RoPE table too small — "
|
| 131 |
+
f"offset={offset}, seq_len={seq_len}, table_size={cos.shape[0]}"
|
| 132 |
+
)
|
| 133 |
+
cos_s = jax.lax.dynamic_slice(cos, (offset, 0), (seq_len, half)).astype(x.dtype)
|
| 134 |
+
sin_s = jax.lax.dynamic_slice(sin, (offset, 0), (seq_len, half)).astype(x.dtype)
|
| 135 |
+
while cos_s.ndim < x.ndim - 1:
|
| 136 |
+
cos_s = cos_s[None]
|
| 137 |
+
sin_s = sin_s[None]
|
| 138 |
+
x1, x2 = x[..., :half], x[..., half:]
|
| 139 |
+
return jnp.concatenate(
|
| 140 |
+
[x1 * cos_s - x2 * sin_s, x1 * sin_s + x2 * cos_s],
|
| 141 |
+
axis=-1,
|
| 142 |
+
)
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
# ---------------------------------------------------------------------------
|
| 146 |
+
# Core kernel: block-tiled native GQA SWA
|
| 147 |
+
# ---------------------------------------------------------------------------
|
| 148 |
+
|
| 149 |
+
def _block_gqa_swa(
|
| 150 |
+
q: jnp.ndarray,
|
| 151 |
+
k: jnp.ndarray,
|
| 152 |
+
v: jnp.ndarray,
|
| 153 |
+
window_size: int,
|
| 154 |
+
block_size: int = TPU_BLOCK_SIZE,
|
| 155 |
+
use_remat: bool = True,
|
| 156 |
+
) -> jnp.ndarray:
|
| 157 |
+
"""
|
| 158 |
+
Block-tiled Native GQA Sliding Window Attention.
|
| 159 |
+
|
| 160 |
+
This function is the replacement for the scan+dynamic_slice approach.
|
| 161 |
+
All tensor shapes inside the map body are STATIC — XLA never sees a
|
| 162 |
+
gather-inside-loop pattern, so the hidden [S,B,H,W,D] buffer cannot form.
|
| 163 |
+
|
| 164 |
+
Parameters
|
| 165 |
+
----------
|
| 166 |
+
q : [B, Hq, S, D] — any float dtype (BF16 in production)
|
| 167 |
+
k : [B, Hkv, S, D] — same dtype
|
| 168 |
+
v : [B, Hkv, S, D] — same dtype
|
| 169 |
+
window_size : causal window W; token t attends to [max(0, t-W+1), t]
|
| 170 |
+
block_size : query block size BLK (tune to 128 for TPU v5e)
|
| 171 |
+
use_remat : wrap map body in jax.checkpoint (recommended for training)
|
| 172 |
+
|
| 173 |
+
Returns
|
| 174 |
+
-------
|
| 175 |
+
[B, Hq, S, D] — same dtype as q
|
| 176 |
+
"""
|
| 177 |
+
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
|
| 178 |
+
raise ValueError(
|
| 179 |
+
f"_block_gqa_swa: expected 4-D inputs, "
|
| 180 |
+
f"got q={q.shape} k={k.shape} v={v.shape}"
|
| 181 |
+
)
|
| 182 |
+
|
| 183 |
+
B, Hq, S, D = map(int, q.shape)
|
| 184 |
+
_B, Hkv, Sk, Dk = map(int, k.shape)
|
| 185 |
+
|
| 186 |
+
if _B != B:
|
| 187 |
+
raise ValueError(f"Batch mismatch: q B={B}, k B={_B}")
|
| 188 |
+
if Dk != D:
|
| 189 |
+
raise ValueError(f"Head-dim mismatch: q D={D}, k D={Dk}")
|
| 190 |
+
if Sk != S:
|
| 191 |
+
raise ValueError(f"Sequence-length mismatch: q S={S}, k S={Sk}")
|
| 192 |
+
if Hq % Hkv != 0:
|
| 193 |
+
raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")
|
| 194 |
+
|
| 195 |
+
G = Hq // Hkv # queries per KV head
|
| 196 |
+
W = int(window_size)
|
| 197 |
+
BLK = int(block_size)
|
| 198 |
+
if BLK <= 0:
|
| 199 |
+
raise ValueError(f"block_size must be > 0, got {BLK}")
|
| 200 |
+
|
| 201 |
+
scale = 1.0 / math.sqrt(float(D))
|
| 202 |
+
|
| 203 |
+
# ── Pad sequence to a multiple of BLK ────────────────────────────────────
|
| 204 |
+
n_blocks = (S + BLK - 1) // BLK
|
| 205 |
+
S_pad = n_blocks * BLK # ≥ S, multiple of BLK
|
| 206 |
+
|
| 207 |
+
# ── Pad q along sequence axis ─────────────────────────────────────────────
|
| 208 |
+
# [B, Hq, S_pad, D]
|
| 209 |
+
# The extra (S_pad - S) tokens are zero-padded and trimmed from output.
|
| 210 |
+
q_pad = jnp.pad(q, ((0, 0), (0, 0), (0, S_pad - S), (0, 0)))
|
| 211 |
+
|
| 212 |
+
# ── Pad k/v on the LEFT by (W-1) for causal alignment ────────────────────
|
| 213 |
+
# [B, Hkv, S_pad + W - 1, D]
|
| 214 |
+
#
|
| 215 |
+
# After padding, for query block starting at global position b*BLK:
|
| 216 |
+
# kv slice starts at offset b*BLK in k_pad
|
| 217 |
+
# kv slice has static length BLK + W - 1
|
| 218 |
+
# It covers original positions [b*BLK - (W-1), b*BLK + BLK - 1]
|
| 219 |
+
# which, after clipping to ≥ 0, is exactly the causal window.
|
| 220 |
+
kv_pad_len = S_pad + W - 1 # total padded kv length (static)
|
| 221 |
+
lpad = W - 1 # left zero-padding width
|
| 222 |
+
|
| 223 |
+
k_pad = jnp.pad(k, ((0, 0), (0, 0), (lpad, S_pad - S), (0, 0))) # [B,Hkv,kv_pad_len,D]
|
| 224 |
+
v_pad = jnp.pad(v, ((0, 0), (0, 0), (lpad, S_pad - S), (0, 0)))
|
| 225 |
+
|
| 226 |
+
# ── Static per-block mask ─────────────────────────────────────────────────
|
| 227 |
+
# Compute a boolean mask of shape [BLK, BLK + W - 1].
|
| 228 |
+
# Entry [q_local, k_local] is True (mask=attend) when:
|
| 229 |
+
# (a) k is not from left-padding → k_local >= W - 1 - (something)
|
| 230 |
+
# We cannot compute the absolute positions statically because they depend
|
| 231 |
+
# on block index b. So we record the RELATIVE offsets and apply the
|
| 232 |
+
# offset inside the map body using only static arithmetic on local indices.
|
| 233 |
+
#
|
| 234 |
+
# q_abs = b*BLK + q_local (q_local in 0..BLK-1)
|
| 235 |
+
# k_abs = b*BLK + k_local - lpad (k_local in 0..BLK+W-2)
|
| 236 |
+
#
|
| 237 |
+
# Attend iff:
|
| 238 |
+
# k_abs >= 0 (not left padding)
|
| 239 |
+
# k_abs <= q_abs (causal)
|
| 240 |
+
# q_abs - k_abs < W (within window)
|
| 241 |
+
#
|
| 242 |
+
# q_abs - k_abs = q_local - k_local + lpad (b cancels out!)
|
| 243 |
+
# k_abs >= 0 ↔ k_local >= lpad - b*BLK (depends on b → handle in body)
|
| 244 |
+
# k_abs <= q_abs ↔ k_local - q_local <= lpad (b cancels out! → STATIC)
|
| 245 |
+
#
|
| 246 |
+
# Only the "k_abs >= 0" condition depends on b (first block only).
|
| 247 |
+
# We handle it cheaply with a dynamic mask inside the body.
|
| 248 |
+
# Everything else is b-independent and can be PRECOMPUTED once.
|
| 249 |
+
|
| 250 |
+
kv_len = BLK + W - 1 # static length of kv slice per block
|
| 251 |
+
|
| 252 |
+
q_local = jnp.arange(BLK, dtype=jnp.int32) # [BLK]
|
| 253 |
+
kv_local = jnp.arange(kv_len, dtype=jnp.int32) # [kv_len]
|
| 254 |
+
|
| 255 |
+
# Relative offset: delta[q, k] = q_local[q] - kv_local[k] + lpad
|
| 256 |
+
# = q_abs - k_abs (b-independent)
|
| 257 |
+
delta = q_local[:, None] - kv_local[None, :] + lpad # [BLK, kv_len]
|
| 258 |
+
|
| 259 |
+
# Static masks (b-independent)
|
| 260 |
+
static_future_mask = delta < 0 # k is in the future (causal) → mask
|
| 261 |
+
static_window_mask = delta >= W # k is too far back (window) → mask
|
| 262 |
+
static_base_mask = static_future_mask | static_window_mask # [BLK, kv_len]
|
| 263 |
+
|
| 264 |
+
# ── Map body ──────────────────────────────────────────────────────────────
|
| 265 |
+
|
| 266 |
+
def process_block(b: jnp.ndarray) -> jnp.ndarray:
|
| 267 |
+
"""
|
| 268 |
+
b : scalar int32 block index in [0, n_blocks)
|
| 269 |
+
|
| 270 |
+
Tensor shapes inside this function are ALL STATIC:
|
| 271 |
+
q_blk : [B, Hq, BLK, D]
|
| 272 |
+
k_blk : [B, Hkv, kv_len, D]
|
| 273 |
+
v_blk : [B, Hkv, kv_len, D]
|
| 274 |
+
scores : [B, Hkv, G, BLK, kv_len]
|
| 275 |
+
probs : same
|
| 276 |
+
|
| 277 |
+
XLA has NO gather-inside-loop here. The dynamic_slice start index
|
| 278 |
+
`b * BLK` is a scalar multiply — XLA lowers this to a simple pointer
|
| 279 |
+
offset, not a buffer materialisation.
|
| 280 |
+
"""
|
| 281 |
+
blk_start = b * BLK # scalar traced int32
|
| 282 |
+
|
| 283 |
+
# Static-shape slices — the KEY difference from the scan approach.
|
| 284 |
+
# XLA sees shapes (B,Hq,BLK,D) and (B,Hkv,kv_len,D) as compile-time
|
| 285 |
+
# constants. It cannot stage all blocks simultaneously because
|
| 286 |
+
# lax.map gives it one block at a time with no output accumulation.
|
| 287 |
+
q_blk = jax.lax.dynamic_slice(q_pad, (0, 0, blk_start, 0), (B, Hq, BLK, D))
|
| 288 |
+
k_blk = jax.lax.dynamic_slice(k_pad, (0, 0, blk_start, 0), (B, Hkv, kv_len, D))
|
| 289 |
+
v_blk = jax.lax.dynamic_slice(v_pad, (0, 0, blk_start, 0), (B, Hkv, kv_len, D))
|
| 290 |
+
|
| 291 |
+
# ── Native GQA reshape (zero-copy view) ──────────────────────────
|
| 292 |
+
# [B, Hq, BLK, D] → [B, Hkv, G, BLK, D]
|
| 293 |
+
q_g = q_blk.reshape(B, Hkv, G, BLK, D)
|
| 294 |
+
|
| 295 |
+
# ── Dot-product scores ────────────────────────────────────────────
|
| 296 |
+
# [B, Hkv, G, BLK, kv_len] ← STATIC, FUSED by XLA
|
| 297 |
+
scores = (
|
| 298 |
+
jnp.einsum(
|
| 299 |
+
"bngqd,bnkd->bngqk",
|
| 300 |
+
q_g.astype(jnp.float32),
|
| 301 |
+
k_blk.astype(jnp.float32),
|
| 302 |
+
)
|
| 303 |
+
* scale
|
| 304 |
+
)
|
| 305 |
+
|
| 306 |
+
# ── Masking ───────────────────────────────────────────────────────
|
| 307 |
+
# (a) b-independent mask (precomputed, fused as constant)
|
| 308 |
+
mask = static_base_mask # [BLK, kv_len]
|
| 309 |
+
|
| 310 |
+
# (b) Left-padding mask: k_abs < 0 ↔ kv_local < lpad - blk_start
|
| 311 |
+
# Only non-trivial for b=0 (the very first block).
|
| 312 |
+
# For all subsequent blocks, lpad - blk_start < 0, so no extra masking.
|
| 313 |
+
# We compute it dynamically but it is a single scalar comparison
|
| 314 |
+
# broadcast — XLA will constant-fold it for b > 0 at runtime.
|
| 315 |
+
leftpad_cutoff = lpad - blk_start # scalar int32 (may be negative)
|
| 316 |
+
leftpad_mask = kv_local[None, :] < leftpad_cutoff # [1, kv_len]
|
| 317 |
+
mask = mask | leftpad_mask # [BLK, kv_len]
|
| 318 |
+
|
| 319 |
+
scores = jnp.where(
|
| 320 |
+
mask[None, None, None, :, :], # [1,1,1,BLK,kv_len]
|
| 321 |
+
jnp.full_like(scores, -1e30),
|
| 322 |
+
scores,
|
| 323 |
+
)
|
| 324 |
+
|
| 325 |
+
# ── Softmax + value aggregation (all FP32) ────────────────────────
|
| 326 |
+
probs = jax.nn.softmax(scores, axis=-1) # [B,Hkv,G,BLK,kv_len]
|
| 327 |
+
out_g = jnp.einsum(
|
| 328 |
+
"bngqk,bnkd->bngqd",
|
| 329 |
+
probs,
|
| 330 |
+
v_blk.astype(jnp.float32),
|
| 331 |
+
) # [B, Hkv, G, BLK, D]
|
| 332 |
+
|
| 333 |
+
# ── Reshape + cast back to input dtype ────────────────────────────
|
| 334 |
+
return out_g.reshape(B, Hq, BLK, D).astype(q.dtype) # [B, Hq, BLK, D]
|
| 335 |
+
|
| 336 |
+
# ── Gradient checkpointing ────────────────────────────────────────────────
|
| 337 |
+
# With use_remat=True: JAX recomputes each block's forward during backward.
|
| 338 |
+
# Granularity is per-block (not per-token), so remat overhead is manageable.
|
| 339 |
+
map_fn = jax.checkpoint(process_block) if use_remat else process_block
|
| 340 |
+
|
| 341 |
+
# ── lax.map over block indices ────────────────────────────────────────────
|
| 342 |
+
# Unlike lax.scan, lax.map does NOT accumulate a carry — it writes each
|
| 343 |
+
# block's output directly. No hidden [n_blocks, B, Hq, BLK, D] stack.
|
| 344 |
+
out_blocks = jax.lax.map(map_fn, jnp.arange(n_blocks, dtype=jnp.int32))
|
| 345 |
+
# out_blocks: [n_blocks, B, Hq, BLK, D]
|
| 346 |
+
|
| 347 |
+
# ── Assemble output ───────────────────────────────────────────────────────
|
| 348 |
+
# Transpose to [B, Hq, n_blocks, BLK, D] then reshape to [B, Hq, S_pad, D].
|
| 349 |
+
# The final slice trims padding tokens. All zero-copy views on TPU.
|
| 350 |
+
out = out_blocks.transpose(1, 2, 0, 3, 4).reshape(B, Hq, S_pad, D)
|
| 351 |
+
return out[:, :, :S, :]
|
| 352 |
+
|
| 353 |
+
|
| 354 |
+
# ---------------------------------------------------------------------------
|
| 355 |
+
# Public API
|
| 356 |
+
# ---------------------------------------------------------------------------
|
| 357 |
+
|
| 358 |
+
def flash_splash_attention(
|
| 359 |
+
q: jnp.ndarray,
|
| 360 |
+
k: jnp.ndarray,
|
| 361 |
+
v: jnp.ndarray,
|
| 362 |
+
window_size: int,
|
| 363 |
+
backend: Optional[str] = None,
|
| 364 |
+
use_gqa: bool = True,
|
| 365 |
+
start_pos: int = 0,
|
| 366 |
+
block_size: int = TPU_BLOCK_SIZE,
|
| 367 |
+
use_remat: bool = True,
|
| 368 |
+
) -> jnp.ndarray:
|
| 369 |
+
"""
|
| 370 |
+
Block-tiled GQA-native Sliding Window Attention — main entry point.
|
| 371 |
+
|
| 372 |
+
Drop-in replacement for the previous flash_splash_attention.
|
| 373 |
+
GQA is always native (no KV expansion). `use_gqa` is accepted for
|
| 374 |
+
API compatibility only.
|
| 375 |
+
|
| 376 |
+
Parameters
|
| 377 |
+
----------
|
| 378 |
+
q, k, v : [B, Hq, S, D] / [B, Hkv, S, D]
|
| 379 |
+
window_size : causal window W
|
| 380 |
+
backend : reserved for future Pallas/Flash kernel dispatch (no-op)
|
| 381 |
+
use_gqa : API-compat flag; GQA is always native here
|
| 382 |
+
start_pos : KV-cache offset for decode mode (training path: leave at 0)
|
| 383 |
+
block_size : query block size BLK (128 recommended for TPU v5e)
|
| 384 |
+
use_remat : gradient checkpointing in map body (recommended for training)
|
| 385 |
+
|
| 386 |
+
Returns
|
| 387 |
+
-------
|
| 388 |
+
[B, Hq, S, D] — same dtype as q
|
| 389 |
+
"""
|
| 390 |
+
_ = backend
|
| 391 |
+
_ = use_gqa
|
| 392 |
+
_ = start_pos # TODO: shift q/kv slicing for decode-mode KV cache
|
| 393 |
+
|
| 394 |
+
return _block_gqa_swa(
|
| 395 |
+
q=q,
|
| 396 |
+
k=k,
|
| 397 |
+
v=v,
|
| 398 |
+
window_size=int(window_size),
|
| 399 |
+
block_size=int(block_size),
|
| 400 |
+
use_remat=use_remat,
|
| 401 |
+
)
|
| 402 |
+
def decode_swa(
|
| 403 |
+
q: jnp.ndarray,
|
| 404 |
+
k: jnp.ndarray,
|
| 405 |
+
v: jnp.ndarray,
|
| 406 |
+
) -> jnp.ndarray:
|
| 407 |
+
"""
|
| 408 |
+
Decode-time SWA for one-token generation.
|
| 409 |
+
|
| 410 |
+
q: [B, Hq, 1, D]
|
| 411 |
+
k: [B, Hkv, W, D]
|
| 412 |
+
v: [B, Hkv, W, D]
|
| 413 |
+
|
| 414 |
+
Returns: [B, Hq, 1, D]
|
| 415 |
+
"""
|
| 416 |
+
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
|
| 417 |
+
raise ValueError(
|
| 418 |
+
f"decode_swa: expected 4-D tensors, got q={q.shape}, k={k.shape}, v={v.shape}"
|
| 419 |
+
)
|
| 420 |
+
|
| 421 |
+
B, Hq, S, D = map(int, q.shape)
|
| 422 |
+
_, Hkv, W, Dk = map(int, k.shape)
|
| 423 |
+
|
| 424 |
+
if S != 1:
|
| 425 |
+
raise ValueError(f"decode_swa expects one token, got S={S}")
|
| 426 |
+
if D != Dk:
|
| 427 |
+
raise ValueError(f"Head-dim mismatch: q D={D}, k D={Dk}")
|
| 428 |
+
if Hq % Hkv != 0:
|
| 429 |
+
raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")
|
| 430 |
+
|
| 431 |
+
G = Hq // Hkv
|
| 432 |
+
scale = 1.0 / math.sqrt(float(D))
|
| 433 |
+
|
| 434 |
+
q_g = q.reshape(B, Hkv, G, 1, D)
|
| 435 |
+
|
| 436 |
+
scores = (
|
| 437 |
+
jnp.einsum(
|
| 438 |
+
"bngqd,bnkd->bngqk",
|
| 439 |
+
q_g.astype(jnp.float32),
|
| 440 |
+
k.astype(jnp.float32),
|
| 441 |
+
)
|
| 442 |
+
* scale
|
| 443 |
+
) # [B, Hkv, G, 1, W]
|
| 444 |
+
|
| 445 |
+
probs = jax.nn.softmax(scores, axis=-1)
|
| 446 |
+
|
| 447 |
+
out_g = jnp.einsum(
|
| 448 |
+
"bngqk,bnkd->bngqd",
|
| 449 |
+
probs,
|
| 450 |
+
v.astype(jnp.float32),
|
| 451 |
+
)
|
| 452 |
+
|
| 453 |
+
return out_g.reshape(B, Hq, 1, D).astype(q.dtype)
|
| 454 |
+
|
| 455 |
+
def local_causal_attention(
|
| 456 |
+
q: jnp.ndarray,
|
| 457 |
+
k: jnp.ndarray,
|
| 458 |
+
v: jnp.ndarray,
|
| 459 |
+
) -> jnp.ndarray:
|
| 460 |
+
"""
|
| 461 |
+
Full causal attention with native GQA — no windowing.
|
| 462 |
+
|
| 463 |
+
⚠ Creates an O(S²) score tensor [B, Hkv, G, S, S].
|
| 464 |
+
Use only for short sequences, unit tests, or reference baselines.
|
| 465 |
+
|
| 466 |
+
q: [B, Hq, S, D] | k, v: [B, Hkv, S, D]
|
| 467 |
+
returns: [B, Hq, S, D]
|
| 468 |
+
"""
|
| 469 |
+
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
|
| 470 |
+
raise ValueError(
|
| 471 |
+
f"local_causal_attention: expected 4-D tensors, "
|
| 472 |
+
f"got q={q.shape} k={k.shape} v={v.shape}"
|
| 473 |
+
)
|
| 474 |
+
B, Hq, S, D = map(int, q.shape)
|
| 475 |
+
_, Hkv, K, _ = map(int, k.shape)
|
| 476 |
+
if Hq % Hkv != 0:
|
| 477 |
+
raise ValueError(f"Hq={Hq} must be divisible by Hkv={Hkv}")
|
| 478 |
+
G = Hq // Hkv
|
| 479 |
+
scale = 1.0 / math.sqrt(float(D))
|
| 480 |
+
q_g = q.reshape(B, Hkv, G, S, D)
|
| 481 |
+
scores = (
|
| 482 |
+
jnp.einsum("bngsd,bnkd->bngsk", q_g.astype(jnp.float32), k.astype(jnp.float32))
|
| 483 |
+
* scale
|
| 484 |
+
)
|
| 485 |
+
qi = jnp.arange(S, dtype=jnp.int32)[:, None]
|
| 486 |
+
ki = jnp.arange(K, dtype=jnp.int32)[None, :]
|
| 487 |
+
scores = jnp.where(ki > qi, -1e30, scores)
|
| 488 |
+
probs = jax.nn.softmax(scores, axis=-1)
|
| 489 |
+
out_g = jnp.einsum("bngsk,bnkd->bngsd", probs, v.astype(jnp.float32))
|
| 490 |
+
return out_g.reshape(B, Hq, S, D).astype(q.dtype)
|
| 491 |
+
|
| 492 |
+
|
| 493 |
+
# ---------------------------------------------------------------------------
|
| 494 |
+
# Validation suite
|
| 495 |
+
# ---------------------------------------------------------------------------
|
| 496 |
+
|
| 497 |
+
if __name__ == "__main__":
|
| 498 |
+
import sys
|
| 499 |
+
|
| 500 |
+
PASS = "\033[92m✓\033[0m"
|
| 501 |
+
FAIL = "\033[91m✗\033[0m"
|
| 502 |
+
HDR = "\033[1;94m"
|
| 503 |
+
RST = "\033[0m"
|
| 504 |
+
failures = 0
|
| 505 |
+
|
| 506 |
+
def section(t): print(f"\n{HDR}{'─'*62}{RST}\n{HDR} {t}{RST}\n{HDR}{'─'*62}{RST}")
|
| 507 |
+
def ok(m): print(f" {PASS} {m}")
|
| 508 |
+
def fail(m):
|
| 509 |
+
global failures; failures += 1
|
| 510 |
+
print(f" {FAIL} {m}", file=sys.stderr)
|
| 511 |
+
|
| 512 |
+
print(f"\n{HDR}{'═'*62}{RST}")
|
| 513 |
+
print(f"{HDR} Veylon Attention — Block-tiled GQA SWA — Validation{RST}")
|
| 514 |
+
print(f"{HDR}{'═'*62}{RST}")
|
| 515 |
+
|
| 516 |
+
# ── 1. Shape and dtype ───────────────────────────────────────────────────
|
| 517 |
+
section("1 · Shape and dtype correctness")
|
| 518 |
+
|
| 519 |
+
B, Hq, Hkv, S, D, W = 1, 8, 2, 128, 64, 32
|
| 520 |
+
ks = jax.random.split(jax.random.PRNGKey(0), 3)
|
| 521 |
+
q = jax.random.normal(ks[0], (B, Hq, S, D), dtype=jnp.bfloat16)
|
| 522 |
+
k = jax.random.normal(ks[1], (B, Hkv, S, D), dtype=jnp.bfloat16)
|
| 523 |
+
v = jax.random.normal(ks[2], (B, Hkv, S, D), dtype=jnp.bfloat16)
|
| 524 |
+
|
| 525 |
+
out = flash_splash_attention(q, k, v, window_size=W, block_size=32)
|
| 526 |
+
|
| 527 |
+
if out.shape == (B, Hq, S, D): ok(f"Output shape : {out.shape}")
|
| 528 |
+
else: fail(f"Shape wrong — expected {(B,Hq,S,D)}, got {out.shape}")
|
| 529 |
+
if out.dtype == jnp.bfloat16: ok(f"Output dtype : {out.dtype}")
|
| 530 |
+
else: fail(f"dtype wrong — expected bfloat16, got {out.dtype}")
|
| 531 |
+
if not jnp.any(jnp.isnan(out)): ok("No NaNs")
|
| 532 |
+
else: fail("Output contains NaNs")
|
| 533 |
+
|
| 534 |
+
# ── 2. Numerical agreement with reference SWA ────────────────────────────
|
| 535 |
+
section("2 · Numerical agreement: block SWA ≈ reference SWA (W=S)")
|
| 536 |
+
|
| 537 |
+
B_, Hq_, Hkv_, S_, D_ = 1, 4, 2, 48, 16
|
| 538 |
+
ks = jax.random.split(jax.random.PRNGKey(1), 3)
|
| 539 |
+
qn = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32)
|
| 540 |
+
kn = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 541 |
+
vn = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 542 |
+
|
| 543 |
+
for W_t in [4, 8, 16, S_]:
|
| 544 |
+
out_blk = flash_splash_attention(qn, kn, vn, window_size=W_t, block_size=16)
|
| 545 |
+
out_ref = local_causal_attention(qn, kn, vn) if W_t == S_ else None
|
| 546 |
+
|
| 547 |
+
# Build reference SWA inline for each W_t
|
| 548 |
+
G_ = Hq_ // Hkv_
|
| 549 |
+
sc = 1.0 / math.sqrt(D_)
|
| 550 |
+
qg = qn.reshape(B_, Hkv_, G_, S_, D_)
|
| 551 |
+
sc_ref = jnp.einsum("bngsd,bnkd->bngsk", qg, kn) * sc
|
| 552 |
+
qi = jnp.arange(S_)[:, None]; ki = jnp.arange(S_)[None, :]
|
| 553 |
+
sc_ref = jnp.where((ki > qi) | (qi - ki >= W_t), -1e30, sc_ref)
|
| 554 |
+
pr_ref = jax.nn.softmax(sc_ref, axis=-1)
|
| 555 |
+
out_swa_ref = jnp.einsum("bngsk,bnkd->bngsd", pr_ref, vn).reshape(B_, Hq_, S_, D_)
|
| 556 |
+
|
| 557 |
+
err = float(jnp.max(jnp.abs(out_blk - out_swa_ref)))
|
| 558 |
+
if err < 1e-4: ok(f"W={W_t:3d} max|block - ref_swa| = {err:.2e}")
|
| 559 |
+
else: fail(f"W={W_t:3d} MISMATCH: max err = {err:.2e}")
|
| 560 |
+
|
| 561 |
+
# ── 3. Memory scaling ────────────────────────────────────────────────────
|
| 562 |
+
section("3 · Memory scaling (2× S → ~2× memory, not 4×)")
|
| 563 |
+
print(" Shape + completion check at S = 256, 512, 1024, 2048")
|
| 564 |
+
for S_t in [256, 512, 1024, 2048]:
|
| 565 |
+
ks = jax.random.split(jax.random.PRNGKey(S_t), 3)
|
| 566 |
+
qt = jax.random.normal(ks[0], (1, 8, S_t, 64), dtype=jnp.bfloat16)
|
| 567 |
+
kt = jax.random.normal(ks[1], (1, 2, S_t, 64), dtype=jnp.bfloat16)
|
| 568 |
+
vt = jax.random.normal(ks[2], (1, 2, S_t, 64), dtype=jnp.bfloat16)
|
| 569 |
+
ot = flash_splash_attention(qt, kt, vt, window_size=128)
|
| 570 |
+
if ot.shape == (1, 8, S_t, 64): ok(f"S={S_t:5d} → {ot.shape}")
|
| 571 |
+
else: fail(f"S={S_t} wrong shape {ot.shape}")
|
| 572 |
+
|
| 573 |
+
# ── 4. Native GQA isolation ──────────────────────────────────────────────
|
| 574 |
+
section("4 �� Native GQA head isolation (no KV duplication)")
|
| 575 |
+
B_, Hq_, Hkv_, S_, D_, W_ = 1, 4, 2, 32, 16, 16
|
| 576 |
+
G_ = Hq_ // Hkv_
|
| 577 |
+
ks = jax.random.split(jax.random.PRNGKey(7), 3)
|
| 578 |
+
qg = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32)
|
| 579 |
+
kgq = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 580 |
+
vgq = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 581 |
+
base = flash_splash_attention(qg, kgq, vgq, window_size=W_, block_size=16)
|
| 582 |
+
zero0 = flash_splash_attention(qg, kgq.at[:,0].set(0.), vgq.at[:,0].set(0.), window_size=W_, block_size=16)
|
| 583 |
+
changed = not jnp.allclose(base[:, :G_], zero0[:, :G_], atol=1e-4)
|
| 584 |
+
unchanged = jnp.allclose( base[:, G_:], zero0[:, G_:], atol=1e-4)
|
| 585 |
+
ok(f"Q heads 0..{G_-1} changed when KV head 0 zeroed") if changed else fail("GQA dependency broken")
|
| 586 |
+
ok(f"Q heads {G_}..{Hq_-1} unchanged (correct isolation)") if unchanged else fail("Cross-group contamination")
|
| 587 |
+
|
| 588 |
+
# ── 5. Causality ─────────────────────────────────────────────────────────
|
| 589 |
+
section("5 · Causality")
|
| 590 |
+
B_, Hq_, Hkv_, S_, D_, W_ = 1, 4, 2, 24, 8, 8
|
| 591 |
+
pivot = S_ // 2
|
| 592 |
+
ks = jax.random.split(jax.random.PRNGKey(42), 3)
|
| 593 |
+
qc = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32)
|
| 594 |
+
kc = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 595 |
+
vc = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 596 |
+
out_base = flash_splash_attention(qc, kc, vc, window_size=W_, block_size=8)
|
| 597 |
+
out_corr = flash_splash_attention(qc,
|
| 598 |
+
kc.at[:,:,pivot:].set(999.),
|
| 599 |
+
vc.at[:,:,pivot:].set(999.),
|
| 600 |
+
window_size=W_, block_size=8)
|
| 601 |
+
if jnp.allclose(out_base[:,:,:pivot], out_corr[:,:,:pivot], atol=1e-5):
|
| 602 |
+
ok(f"Tokens 0..{pivot-1} unaffected by corruption of tokens {pivot}+")
|
| 603 |
+
else:
|
| 604 |
+
fail("Causality violated — future tokens leaked into past outputs")
|
| 605 |
+
|
| 606 |
+
# ── 6. Window boundary ───────────────────────────────────────────────────
|
| 607 |
+
section("6 · Window boundary (no leakage past W tokens)")
|
| 608 |
+
B_, Hq_, Hkv_, S_, D_, W_ = 1, 2, 1, 24, 8, 4
|
| 609 |
+
ks = jax.random.split(jax.random.PRNGKey(11), 3)
|
| 610 |
+
qw = jax.random.normal(ks[0], (B_, Hq_, S_, D_), dtype=jnp.float32)
|
| 611 |
+
kw = jax.random.normal(ks[1], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 612 |
+
vw = jax.random.normal(ks[2], (B_, Hkv_, S_, D_), dtype=jnp.float32)
|
| 613 |
+
out_w = flash_splash_attention(qw, kw, vw, window_size=W_, block_size=4)
|
| 614 |
+
out_wm = flash_splash_attention(qw,
|
| 615 |
+
kw.at[:,:,0].set(999.),
|
| 616 |
+
vw.at[:,:,0].set(999.),
|
| 617 |
+
window_size=W_, block_size=4)
|
| 618 |
+
pivot_w = W_ # first token where token 0 is outside the window
|
| 619 |
+
if jnp.allclose(out_w[:,:,pivot_w:], out_wm[:,:,pivot_w:], atol=1e-5):
|
| 620 |
+
ok(f"Tokens {pivot_w}+ unaffected by modifying token 0 (W={W_} boundary)")
|
| 621 |
+
else:
|
| 622 |
+
fail(f"Window boundary violated — token 0 leaked into token {pivot_w}+")
|
| 623 |
+
|
| 624 |
+
# ── 7. Non-power-of-2 sequence length ────────────────────────────────────
|
| 625 |
+
section("7 · Non-power-of-2 sequence lengths (S=100, 200, 500)")
|
| 626 |
+
for S_t in [100, 200, 500]:
|
| 627 |
+
ks = jax.random.split(jax.random.PRNGKey(S_t+1), 3)
|
| 628 |
+
qt = jax.random.normal(ks[0], (1, 4, S_t, 16), dtype=jnp.float32)
|
| 629 |
+
kt = jax.random.normal(ks[1], (1, 2, S_t, 16), dtype=jnp.float32)
|
| 630 |
+
vt = jax.random.normal(ks[2], (1, 2, S_t, 16), dtype=jnp.float32)
|
| 631 |
+
ot = flash_splash_attention(qt, kt, vt, window_size=32, block_size=32)
|
| 632 |
+
if ot.shape == (1, 4, S_t, 16): ok(f"S={S_t} → {ot.shape}")
|
| 633 |
+
else: fail(f"S={S_t} wrong shape {ot.shape}")
|
| 634 |
+
|
| 635 |
+
# ── Summary ──────────────────────────────────────────────────────────────
|
| 636 |
+
print(f"\n{HDR}{'═'*62}{RST}")
|
| 637 |
+
if failures == 0: print(f" {PASS} All tests passed.")
|
| 638 |
+
else: print(f" {FAIL} {failures} test(s) failed.", file=sys.stderr)
|
| 639 |
+
print(f"{HDR}{'═'*62}{RST}\n")
|
| 640 |
+
sys.exit(failures)
|
veylon_model.py
ADDED
|
@@ -0,0 +1,383 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import numpy as np
|
| 4 |
+
import keras
|
| 5 |
+
from keras import layers, ops
|
| 6 |
+
import jax
|
| 7 |
+
from veylon_attention import flash_splash_attention , decode_swa
|
| 8 |
+
|
| 9 |
+
try:
|
| 10 |
+
from config import (
|
| 11 |
+
CONTEXT,
|
| 12 |
+
vocab_size as Vocab_size,
|
| 13 |
+
D_MODEL,
|
| 14 |
+
numberoflayers,
|
| 15 |
+
numberofheads,
|
| 16 |
+
d_Latent,
|
| 17 |
+
ffn_mult,
|
| 18 |
+
swa_window,
|
| 19 |
+
num_kv_heads,
|
| 20 |
+
use_moe,
|
| 21 |
+
moe_num_experts,
|
| 22 |
+
moe_top_k,
|
| 23 |
+
)
|
| 24 |
+
except Exception:
|
| 25 |
+
CONTEXT = 2048
|
| 26 |
+
Vocab_size = 32000
|
| 27 |
+
D_MODEL = 512
|
| 28 |
+
numberoflayers = 8
|
| 29 |
+
numberofheads = 8
|
| 30 |
+
d_Latent = 128
|
| 31 |
+
ffn_mult = 3.5
|
| 32 |
+
swa_window = 1024
|
| 33 |
+
num_kv_heads = 2
|
| 34 |
+
use_moe = False
|
| 35 |
+
moe_num_experts = 8
|
| 36 |
+
moe_top_k = 2
|
| 37 |
+
|
| 38 |
+
keras.mixed_precision.set_global_policy('mixed_bfloat16')
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
@keras.saving.register_keras_serializable()
|
| 42 |
+
class RMSNorm(layers.Layer):
|
| 43 |
+
def __init__(self, epsilon=1e-5, **kwargs):
|
| 44 |
+
super().__init__(**kwargs)
|
| 45 |
+
self.epsilon = epsilon
|
| 46 |
+
|
| 47 |
+
def build(self, input_shape):
|
| 48 |
+
self.weight = self.add_weight(shape=(input_shape[-1],), initializer='ones', name='gamma')
|
| 49 |
+
|
| 50 |
+
def call(self, x):
|
| 51 |
+
x_fp32 = ops.cast(x, 'float32')
|
| 52 |
+
rms = ops.sqrt(ops.mean(ops.square(x_fp32), axis=-1, keepdims=True) + self.epsilon)
|
| 53 |
+
out = x_fp32 / rms
|
| 54 |
+
return ops.cast(out, x.dtype) * self.weight
|
| 55 |
+
|
| 56 |
+
def get_config(self):
|
| 57 |
+
cfg = super().get_config()
|
| 58 |
+
cfg.update({'epsilon': self.epsilon})
|
| 59 |
+
return cfg
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
@keras.saving.register_keras_serializable()
|
| 63 |
+
class RotaryEmbedding(layers.Layer):
|
| 64 |
+
def __init__(self, max_seq_len, dim, theta=10000.0, **kwargs):
|
| 65 |
+
super().__init__(**kwargs)
|
| 66 |
+
if dim % 2 != 0:
|
| 67 |
+
raise ValueError('RotaryEmbedding dim must be even.')
|
| 68 |
+
self.max_seq_len = max_seq_len
|
| 69 |
+
self.dim = dim
|
| 70 |
+
self.theta = theta
|
| 71 |
+
|
| 72 |
+
def build(self, input_shape):
|
| 73 |
+
half = self.dim // 2
|
| 74 |
+
inv_freq = 1.0 / (self.theta ** (np.arange(0, self.dim, 2).astype(np.float32) / self.dim))
|
| 75 |
+
positions = np.arange(self.max_seq_len, dtype=np.float32)
|
| 76 |
+
freqs = positions[:, None] * inv_freq[None, :]
|
| 77 |
+
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')
|
| 78 |
+
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')
|
| 79 |
+
|
| 80 |
+
def call(self, x, offset=0):
|
| 81 |
+
seq_len = x.shape[1]
|
| 82 |
+
if seq_len is None:
|
| 83 |
+
seq_len = ops.shape(x)[1]
|
| 84 |
+
half = self.dim // 2
|
| 85 |
+
if offset + seq_len > self.max_seq_len:
|
| 86 |
+
raise ValueError(f'RoPE table too small: offset={offset}, seq_len={seq_len}, max={self.max_seq_len}')
|
| 87 |
+
cos = self.cos[offset:offset + seq_len, :half]
|
| 88 |
+
sin = self.sin[offset:offset + seq_len, :half]
|
| 89 |
+
cos = ops.cast(ops.reshape(cos, (1, seq_len, 1, half)), x.dtype)
|
| 90 |
+
sin = ops.cast(ops.reshape(sin, (1, seq_len, 1, half)), x.dtype)
|
| 91 |
+
x1 = x[..., :half]
|
| 92 |
+
x2 = x[..., half:]
|
| 93 |
+
return ops.concatenate([x1 * cos - x2 * sin, x1 * sin + x2 * cos], axis=-1)
|
| 94 |
+
|
| 95 |
+
def get_config(self):
|
| 96 |
+
cfg = super().get_config()
|
| 97 |
+
cfg.update({'max_seq_len': self.max_seq_len, 'dim': self.dim, 'theta': self.theta})
|
| 98 |
+
return cfg
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
@keras.saving.register_keras_serializable()
|
| 102 |
+
class SwiGLUFFN(layers.Layer):
|
| 103 |
+
def __init__(self, d_model, hidden_mult=3.5, **kwargs):
|
| 104 |
+
super().__init__(**kwargs)
|
| 105 |
+
self.d_model_arg = d_model
|
| 106 |
+
self.hidden_mult = hidden_mult
|
| 107 |
+
self.hidden_dim = int(d_model * hidden_mult * 2 / 3)
|
| 108 |
+
self.hidden_dim = ((self.hidden_dim + 63) // 64) * 64
|
| 109 |
+
|
| 110 |
+
def build(self, input_shape):
|
| 111 |
+
d_model = input_shape[-1]
|
| 112 |
+
self.gate_up_proj = self.add_weight(shape=(d_model, 2 * self.hidden_dim), initializer='glorot_uniform', name='gate_up_proj')
|
| 113 |
+
self.down_proj = self.add_weight(shape=(self.hidden_dim, d_model), initializer='glorot_uniform', name='down_proj')
|
| 114 |
+
|
| 115 |
+
def call(self, x, training=False):
|
| 116 |
+
gate_up = ops.matmul(x, self.gate_up_proj)
|
| 117 |
+
gate, up = ops.split(gate_up, 2, axis=-1)
|
| 118 |
+
return ops.matmul(ops.silu(gate) * up, self.down_proj)
|
| 119 |
+
|
| 120 |
+
def get_config(self):
|
| 121 |
+
cfg = super().get_config()
|
| 122 |
+
cfg.update({'d_model': self.d_model_arg, 'hidden_mult': self.hidden_mult})
|
| 123 |
+
return cfg
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
@keras.saving.register_keras_serializable()
|
| 127 |
+
class MoE_FFN(layers.Layer):
|
| 128 |
+
def __init__(self, d_model, num_experts=8, top_k=2, hidden_mult=3.5, **kwargs):
|
| 129 |
+
super().__init__(**kwargs)
|
| 130 |
+
self.d_model_arg = d_model
|
| 131 |
+
self.num_experts = num_experts
|
| 132 |
+
self.top_k = top_k
|
| 133 |
+
self.hidden_mult = hidden_mult
|
| 134 |
+
self.experts = [SwiGLUFFN(d_model, hidden_mult) for _ in range(num_experts)]
|
| 135 |
+
self.router = layers.Dense(num_experts, use_bias=False)
|
| 136 |
+
|
| 137 |
+
def _load_balancing_loss(self, router_logits, top_k_indices):
|
| 138 |
+
router_probs = ops.softmax(router_logits, axis=-1)
|
| 139 |
+
mask = ops.one_hot(top_k_indices, self.num_experts)
|
| 140 |
+
mask = ops.sum(mask, axis=2)
|
| 141 |
+
f = ops.mean(mask, axis=(0, 1))
|
| 142 |
+
p = ops.mean(router_probs, axis=(0, 1))
|
| 143 |
+
return self.num_experts * ops.sum(f * p)
|
| 144 |
+
|
| 145 |
+
def call(self, x, training=False):
|
| 146 |
+
router_logits = self.router(x)
|
| 147 |
+
top_logits, top_idx = ops.top_k(router_logits, self.top_k)
|
| 148 |
+
top_weights = ops.softmax(top_logits, axis=-1)
|
| 149 |
+
if training:
|
| 150 |
+
self.add_loss(self._load_balancing_loss(router_logits, top_idx))
|
| 151 |
+
out = ops.zeros_like(x)
|
| 152 |
+
for kk in range(self.top_k):
|
| 153 |
+
idx = top_idx[..., kk]
|
| 154 |
+
w = top_weights[..., kk]
|
| 155 |
+
for e in range(self.num_experts):
|
| 156 |
+
mask = ops.cast(idx == e, x.dtype)
|
| 157 |
+
expert_out = self.experts[e](x * mask[..., None], training=training)
|
| 158 |
+
out += expert_out * w[..., None] * mask[..., None]
|
| 159 |
+
return out
|
| 160 |
+
|
| 161 |
+
def get_config(self):
|
| 162 |
+
cfg = super().get_config()
|
| 163 |
+
cfg.update({'d_model': self.d_model_arg, 'num_experts': self.num_experts, 'top_k': self.top_k, 'hidden_mult': self.hidden_mult})
|
| 164 |
+
return cfg
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
@keras.saving.register_keras_serializable()
|
| 168 |
+
class MLAttention(layers.Layer):
|
| 169 |
+
def __init__(self, d_model, n_heads, d_latent, max_seq_len, num_kv_heads=2, swa_window=1024, attn_dropout=0.0, **kwargs):
|
| 170 |
+
super().__init__(**kwargs)
|
| 171 |
+
if d_model % n_heads != 0:
|
| 172 |
+
raise ValueError('d_model must be divisible by n_heads')
|
| 173 |
+
if n_heads % num_kv_heads != 0:
|
| 174 |
+
raise ValueError('n_heads must be divisible by num_kv_heads')
|
| 175 |
+
self.d_model = d_model
|
| 176 |
+
self.n_heads = n_heads
|
| 177 |
+
self.num_kv_heads = num_kv_heads
|
| 178 |
+
self.group_size = n_heads // num_kv_heads
|
| 179 |
+
self.d_head = d_model // n_heads
|
| 180 |
+
self.d_latent = d_latent
|
| 181 |
+
self.max_seq_len = max_seq_len
|
| 182 |
+
self.swa_window = swa_window
|
| 183 |
+
self.dropout = layers.Dropout(attn_dropout)
|
| 184 |
+
|
| 185 |
+
def build(self, input_shape):
|
| 186 |
+
self.W_qc = self.add_weight(shape=(self.d_model, self.d_model + self.d_latent), initializer='glorot_uniform', name='W_qc')
|
| 187 |
+
self.W_kv = self.add_weight(shape=(self.d_latent, self.num_kv_heads * 2 * self.d_head), initializer='glorot_uniform', name='W_kv')
|
| 188 |
+
self.W_o = self.add_weight(shape=(self.d_model, self.d_model), initializer='glorot_uniform', name='Wo')
|
| 189 |
+
self.rope = RotaryEmbedding(self.max_seq_len, self.d_head)
|
| 190 |
+
|
| 191 |
+
def _project_kv(self, c):
|
| 192 |
+
kv = ops.matmul(c, self.W_kv)
|
| 193 |
+
return ops.split(kv, 2, axis=-1)
|
| 194 |
+
|
| 195 |
+
def call(self, x, training=False):
|
| 196 |
+
B = ops.shape(x)[0]
|
| 197 |
+
S = ops.shape(x)[1]
|
| 198 |
+
qc = ops.matmul(x, self.W_qc)
|
| 199 |
+
q_proj, c = ops.split(qc, [self.d_model], axis=-1)
|
| 200 |
+
q = ops.reshape(q_proj, (B, S, self.n_heads, self.d_head))
|
| 201 |
+
q = self.rope(q, offset=0)
|
| 202 |
+
q = ops.transpose(q, (0, 2, 1, 3))
|
| 203 |
+
k, v = self._project_kv(c)
|
| 204 |
+
k = ops.reshape(k, (B, S, self.num_kv_heads, self.d_head))
|
| 205 |
+
v = ops.reshape(v, (B, S, self.num_kv_heads, self.d_head))
|
| 206 |
+
k = self.rope(k, offset=0)
|
| 207 |
+
k = ops.transpose(k, (0, 2, 1, 3))
|
| 208 |
+
v = ops.transpose(v, (0, 2, 1, 3))
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
out = flash_splash_attention(
|
| 212 |
+
q,
|
| 213 |
+
k,
|
| 214 |
+
v,
|
| 215 |
+
window_size=min(self.swa_window, self.max_seq_len),
|
| 216 |
+
backend=jax.default_backend(),
|
| 217 |
+
use_gqa=True,
|
| 218 |
+
)
|
| 219 |
+
|
| 220 |
+
out = ops.transpose(out, (0, 2, 1, 3))
|
| 221 |
+
out = ops.reshape(out, (B, S, self.d_model))
|
| 222 |
+
out = self.dropout(out, training=training)
|
| 223 |
+
return ops.matmul(out, self.W_o)
|
| 224 |
+
|
| 225 |
+
def generate_step(self, x, cache_k=None, cache_v=None, cache_pos=0):
|
| 226 |
+
B = ops.shape(x)[0]
|
| 227 |
+
S = ops.shape(x)[1]
|
| 228 |
+
|
| 229 |
+
qc = ops.matmul(x, self.W_qc)
|
| 230 |
+
q_proj, c = ops.split(qc, [self.d_model], axis=-1)
|
| 231 |
+
|
| 232 |
+
q = ops.reshape(q_proj, (B, S, self.n_heads, self.d_head))
|
| 233 |
+
q = self.rope(q, offset=cache_pos)
|
| 234 |
+
q = ops.transpose(q, (0, 2, 1, 3))
|
| 235 |
+
|
| 236 |
+
k, v = self._project_kv(c)
|
| 237 |
+
k = ops.reshape(k, (B, S, self.num_kv_heads, self.d_head))
|
| 238 |
+
v = ops.reshape(v, (B, S, self.num_kv_heads, self.d_head))
|
| 239 |
+
k = self.rope(k, offset=cache_pos)
|
| 240 |
+
k = ops.transpose(k, (0, 2, 1, 3))
|
| 241 |
+
v = ops.transpose(v, (0, 2, 1, 3))
|
| 242 |
+
|
| 243 |
+
# Prefill path: no previous cache
|
| 244 |
+
if cache_k is None:
|
| 245 |
+
out = flash_splash_attention(
|
| 246 |
+
q,
|
| 247 |
+
k,
|
| 248 |
+
v,
|
| 249 |
+
window_size=min(self.swa_window, self.max_seq_len),
|
| 250 |
+
backend=jax.default_backend(),
|
| 251 |
+
use_gqa=True,
|
| 252 |
+
)
|
| 253 |
+
new_k = k[:, :, -self.swa_window :, :]
|
| 254 |
+
new_v = v[:, :, -self.swa_window :, :]
|
| 255 |
+
else:
|
| 256 |
+
# Decode path: one-token generation only
|
| 257 |
+
if S != 1:
|
| 258 |
+
raise ValueError(
|
| 259 |
+
f"generate_step with cache expects S=1, got S={S}"
|
| 260 |
+
)
|
| 261 |
+
|
| 262 |
+
k = ops.concatenate([cache_k, k], axis=2)
|
| 263 |
+
v = ops.concatenate([cache_v, v], axis=2)
|
| 264 |
+
|
| 265 |
+
k = k[:, :, -self.swa_window :, :]
|
| 266 |
+
v = v[:, :, -self.swa_window :, :]
|
| 267 |
+
|
| 268 |
+
out = decode_swa(q, k, v)
|
| 269 |
+
new_k = k
|
| 270 |
+
new_v = v
|
| 271 |
+
|
| 272 |
+
out = ops.transpose(out, (0, 2, 1, 3))
|
| 273 |
+
out = ops.reshape(out, (B, S, self.d_model))
|
| 274 |
+
out = ops.matmul(out, self.W_o)
|
| 275 |
+
|
| 276 |
+
return out, new_k, new_v
|
| 277 |
+
|
| 278 |
+
def get_config(self):
|
| 279 |
+
cfg = super().get_config()
|
| 280 |
+
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})
|
| 281 |
+
return cfg
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
@keras.saving.register_keras_serializable()
|
| 285 |
+
class TransformerBlock(layers.Layer):
|
| 286 |
+
def __init__(self, d_model, n_heads, d_latent, ffn_layer, max_seq_len, num_kv_heads=2, swa_window=1024, **kwargs):
|
| 287 |
+
super().__init__(**kwargs)
|
| 288 |
+
self.d_model = d_model
|
| 289 |
+
self.n_heads = n_heads
|
| 290 |
+
self.d_latent = d_latent
|
| 291 |
+
self.max_seq_len = max_seq_len
|
| 292 |
+
self.num_kv_heads = num_kv_heads
|
| 293 |
+
self.swa_window = swa_window
|
| 294 |
+
self.ffn = keras.saving.deserialize_keras_object(ffn_layer) if isinstance(ffn_layer, dict) else ffn_layer
|
| 295 |
+
self.norm1 = RMSNorm()
|
| 296 |
+
self.norm2 = RMSNorm()
|
| 297 |
+
self.attn = MLAttention(d_model, n_heads, d_latent, max_seq_len, num_kv_heads=num_kv_heads, swa_window=swa_window)
|
| 298 |
+
|
| 299 |
+
def call(self, x, training=False):
|
| 300 |
+
x = x + self.attn(self.norm1(x), training=training)
|
| 301 |
+
x = x + self.ffn(self.norm2(x), training=training)
|
| 302 |
+
return x
|
| 303 |
+
|
| 304 |
+
def generate_step(self, x, cache_k=None, cache_v=None, cache_pos=0):
|
| 305 |
+
attn_out, nck, ncv = self.attn.generate_step(
|
| 306 |
+
self.norm1(x),
|
| 307 |
+
cache_k=cache_k,
|
| 308 |
+
cache_v=cache_v,
|
| 309 |
+
cache_pos=cache_pos,
|
| 310 |
+
)
|
| 311 |
+
x = x + attn_out
|
| 312 |
+
x = x + self.ffn(self.norm2(x), training=False)
|
| 313 |
+
return x, nck, ncv
|
| 314 |
+
def get_config(self):
|
| 315 |
+
cfg = super().get_config()
|
| 316 |
+
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})
|
| 317 |
+
return cfg
|
| 318 |
+
|
| 319 |
+
|
| 320 |
+
@keras.saving.register_keras_serializable()
|
| 321 |
+
class VeylonModel(keras.Model):
|
| 322 |
+
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):
|
| 323 |
+
super().__init__(**kwargs)
|
| 324 |
+
self.vocab_size = vocab_size
|
| 325 |
+
self.d_model = d_model
|
| 326 |
+
self.n_layers = n_layers
|
| 327 |
+
self.n_heads = n_heads
|
| 328 |
+
self.d_latent = d_latent
|
| 329 |
+
self.ffn_mult = ffn_mult
|
| 330 |
+
self.max_seq_len = max_seq_len
|
| 331 |
+
self.use_moe = use_moe
|
| 332 |
+
self.moe_num_experts = moe_num_experts
|
| 333 |
+
self.moe_top_k = moe_top_k
|
| 334 |
+
self.num_kv_heads = num_kv_heads
|
| 335 |
+
self.swa_window = swa_window
|
| 336 |
+
self.embedding = layers.Embedding(vocab_size, d_model, name='token_embedding')
|
| 337 |
+
self.blocks = []
|
| 338 |
+
for i in range(n_layers):
|
| 339 |
+
ffn = MoE_FFN(d_model, moe_num_experts, moe_top_k, ffn_mult) if use_moe else SwiGLUFFN(d_model, ffn_mult)
|
| 340 |
+
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}'))
|
| 341 |
+
self.norm = RMSNorm()
|
| 342 |
+
|
| 343 |
+
def call(self, inputs, training=False):
|
| 344 |
+
x = self.embedding(inputs)
|
| 345 |
+
for block in self.blocks:
|
| 346 |
+
x = block(x, training=training)
|
| 347 |
+
x = self.norm(x)
|
| 348 |
+
embedding_weights = self.embedding.weights[0]
|
| 349 |
+
logits = ops.matmul(x, ops.transpose(embedding_weights))
|
| 350 |
+
return ops.cast(logits, 'float32')
|
| 351 |
+
|
| 352 |
+
def generate_step(self, inputs, cache_k=None, cache_v=None, cache_pos=0):
|
| 353 |
+
x = self.embedding(inputs)
|
| 354 |
+
new_cache_k = []
|
| 355 |
+
new_cache_v = []
|
| 356 |
+
|
| 357 |
+
if cache_k is None:
|
| 358 |
+
cache_k = [None] * len(self.blocks)
|
| 359 |
+
cache_v = [None] * len(self.blocks)
|
| 360 |
+
|
| 361 |
+
for i, block in enumerate(self.blocks):
|
| 362 |
+
x, nck, ncv = block.generate_step(
|
| 363 |
+
x,
|
| 364 |
+
cache_k=cache_k[i],
|
| 365 |
+
cache_v=cache_v[i],
|
| 366 |
+
cache_pos=cache_pos,
|
| 367 |
+
)
|
| 368 |
+
new_cache_k.append(nck)
|
| 369 |
+
new_cache_v.append(ncv)
|
| 370 |
+
|
| 371 |
+
x = self.norm(x)
|
| 372 |
+
embedding_weights = self.embedding.weights[0]
|
| 373 |
+
logits = ops.matmul(x, ops.transpose(embedding_weights))
|
| 374 |
+
logits = ops.cast(logits, 'float32')
|
| 375 |
+
return logits, new_cache_k, new_cache_v
|
| 376 |
+
def get_config(self):
|
| 377 |
+
cfg = super().get_config()
|
| 378 |
+
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})
|
| 379 |
+
return cfg
|
| 380 |
+
|
| 381 |
+
|
| 382 |
+
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):
|
| 383 |
+
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)
|