Arush kumar commited on
Commit
54ad1e5
·
1 Parent(s): a098b7a

Upload 14 files

Browse files
Files changed (13) hide show
  1. check_vocab.py +4 -0
  2. config.py +25 -0
  3. corpus.txt +1 -0
  4. inference.py +193 -0
  5. inspect_weights.py +7 -0
  6. model_build.py +56 -0
  7. token_train.py +20 -0
  8. tokenizer.model +3 -0
  9. tokenizer.vocab +0 -0
  10. train.py +439 -0
  11. train2.py +547 -0
  12. veylon_attention.py +640 -0
  13. 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)