amuzetnoM commited on
Commit
8a518d8
·
verified ·
1 Parent(s): 8035b6f

WYRM v27 FINAL notebook — training on Kaggle T4 x2

Browse files
Files changed (1) hide show
  1. training/wyrm_notebook_v27_FINAL.py +1582 -0
training/wyrm_notebook_v27_FINAL.py ADDED
@@ -0,0 +1,1582 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ WYRM 500M — Kaggle Training Notebook v28 (PRODUCTION)
4
+ ==========================================
5
+
6
+ Self-contained. No patching. No hacks. This IS the training script.
7
+
8
+ v28 PRODUCTION:
9
+ 1. FULL float32 precision — no fp16, no autocast, no GradScaler
10
+ 2. Batch size 4, grad_accum 4 = effective 16 (VRAM allows it: 5.4/15.9 GB)
11
+ 3. Gradient checkpointing KEPT (saves VRAM without cutting precision)
12
+ 4. MANDATORY torch 2.1.2+cu121 on Kaggle (P100 compat, auto-installs Python 3.11)
13
+ 5. numpy<2 installed (torch 2.1.2 compiled against numpy 1.x, Kaggle has 2.4.4)
14
+ 6. BPE tokenizer: gladius_bpe_16k.model (copied from dataset to working dir)
15
+ 7. 15000 total steps, checkpoint every 250, session chaining ready
16
+ 8. Stability > speed. Correct results > throughput.
17
+
18
+ Architecture (WYRM 627M):
19
+ Full: 1024d / 24L / 32H / 32hd / 4096 FFN / seq=1024
20
+ Dummy: 64d / 2L / 4H / 16hd / 256 FFN / seq=128 (same code path)
21
+
22
+ Usage:
23
+ Kaggle: Just run the notebook (auto-detects Kaggle environment)
24
+ Local: python wyrm_notebook.py --dummy # 60-second test run
25
+ Local: python wyrm_notebook.py --local-test # with local paths
26
+ Local: WYRM_DUMMY=1 python wyrm_notebook.py # env var dummy mode
27
+
28
+ Author: Ava Shakil | Artifact Virtual
29
+ Date: 2026-04-08 (Day 58 — v27 FINAL)
30
+ """
31
+
32
+ import os
33
+ import sys
34
+ import json
35
+ import time
36
+ import math
37
+ import shutil
38
+ import random
39
+ import signal
40
+ import logging
41
+ import datetime
42
+ import traceback
43
+ import dataclasses
44
+ from pathlib import Path
45
+ from typing import Dict, List, Optional, Tuple, Any
46
+ from collections import defaultdict
47
+
48
+ # ═══════════════════════════════════════════════════════════════
49
+ # 0. GPU GATE — fail fast if no GPU
50
+ # ═══════════════════════════════════════════════════════════════
51
+
52
+ print("=" * 60)
53
+ print("WYRM 500M — Kaggle Training v27 FINAL")
54
+ print("=" * 60)
55
+
56
+ # ═══════════════════════════════════════════════════════════════
57
+ # MANDATORY: Ensure compatible PyTorch on Kaggle
58
+ # ═══════════════════════════════════════════════════════════════
59
+ # Kaggle ships PyTorch 2.10.0+cu128 on Python 3.12.
60
+ # P100 (sm_60) support was dropped after PyTorch 2.1.x.
61
+ # torch 2.1.2+cu121 only has cp38-cp311 wheels (no cp312).
62
+ #
63
+ # Strategy for P100:
64
+ # 1. Install Python 3.11 via conda (already in Kaggle base image)
65
+ # 2. Install torch 2.1.2+cu121 under Python 3.11
66
+ # 3. Re-exec this script under Python 3.11
67
+ #
68
+ # Strategy for T4+: use stock PyTorch (2.10.0+cu128 supports sm_75+)
69
+ IS_KAGGLE_ENV = os.path.exists('/kaggle')
70
+ _P100_REEXEC = os.environ.get('_WYRM_P100_REEXEC', '')
71
+
72
+ if IS_KAGGLE_ENV and not _P100_REEXEC:
73
+ import subprocess
74
+ try:
75
+ _gpu_check = subprocess.run(
76
+ ['nvidia-smi', '--query-gpu=name', '--format=csv,noheader'],
77
+ capture_output=True, text=True, timeout=10
78
+ )
79
+ _gpu_name = _gpu_check.stdout.strip()
80
+ print(f"GPU: {_gpu_name}")
81
+ if 'P100' in _gpu_name or 'p100' in _gpu_name.lower():
82
+ print("P100 detected — need torch 2.1.2+cu121 on Python 3.11")
83
+ print("Installing Python 3.11 + compatible PyTorch...")
84
+
85
+ # Install Python 3.11 via apt (conda not available on Kaggle)
86
+ subprocess.run(['apt-get', 'update', '-qq'], capture_output=True, timeout=120)
87
+ _r = subprocess.run(
88
+ ['apt-get', 'install', '-y', '-qq', 'python3.11', 'python3.11-venv', 'python3.11-dev'],
89
+ capture_output=True, text=True, timeout=180
90
+ )
91
+ if _r.returncode != 0:
92
+ # python3.11 might not be in default repos, try deadsnakes PPA
93
+ print(" Default repos failed, trying deadsnakes PPA...")
94
+ subprocess.run(
95
+ ['apt-get', 'install', '-y', '-qq', 'software-properties-common'],
96
+ capture_output=True, timeout=60
97
+ )
98
+ subprocess.run(
99
+ ['add-apt-repository', '-y', 'ppa:deadsnakes/ppa'],
100
+ capture_output=True, timeout=60
101
+ )
102
+ subprocess.run(['apt-get', 'update', '-qq'], capture_output=True, timeout=120)
103
+ _r = subprocess.run(
104
+ ['apt-get', 'install', '-y', '-qq', 'python3.11', 'python3.11-venv', 'python3.11-dev'],
105
+ capture_output=True, text=True, timeout=180
106
+ )
107
+ if _r.returncode != 0:
108
+ # Last resort: download wheel directly and use current Python
109
+ print(f" apt install failed: {_r.stderr[-300:]}")
110
+ print(" Trying direct wheel download...")
111
+ # Get Python 3.12 compatible torch built from source with sm_60
112
+ _r = subprocess.run(
113
+ [sys.executable, '-m', 'pip', 'install', '-q', '--force-reinstall',
114
+ 'torch==2.5.1', '--index-url', 'https://download.pytorch.org/whl/cu121'],
115
+ capture_output=True, text=True, timeout=600
116
+ )
117
+ if _r.returncode == 0:
118
+ # Test if 2.5.1+cu121 happens to have sm_60
119
+ _t = subprocess.run(
120
+ [sys.executable, '-c', 'import torch; t=torch.zeros(1,device="cuda"); print("OK")'],
121
+ capture_output=True, text=True, timeout=30
122
+ )
123
+ if 'OK' in _t.stdout:
124
+ print(" torch 2.5.1+cu121 works on P100! ✅")
125
+ os.environ['_WYRM_P100_REEXEC'] = '1'
126
+ os.execv(sys.executable, [sys.executable] + sys.argv)
127
+ print("FATAL: Cannot get P100-compatible PyTorch.")
128
+ sys.exit(1)
129
+
130
+ _py311 = '/usr/bin/python3.11'
131
+ if not os.path.exists(_py311):
132
+ print(f"FATAL: {_py311} not found after install")
133
+ sys.exit(1)
134
+
135
+ print(f" Python 3.11 at: {_py311}")
136
+
137
+ # Install pip for Python 3.11
138
+ subprocess.run(
139
+ [_py311, '-m', 'ensurepip', '--upgrade'],
140
+ capture_output=True, text=True, timeout=60
141
+ )
142
+
143
+ # Install torch 2.1.2+cu121 (from PyTorch index) then sentencepiece (from PyPI)
144
+ print(" Installing torch 2.1.2+cu121...")
145
+ _r = subprocess.run(
146
+ [_py311, '-m', 'pip', 'install', '-q',
147
+ 'torch==2.1.2',
148
+ '--index-url', 'https://download.pytorch.org/whl/cu121'],
149
+ capture_output=True, text=True, timeout=600
150
+ )
151
+ if _r.returncode != 0:
152
+ print(f" torch install failed: {_r.stderr[-500:]}")
153
+ sys.exit(1)
154
+ print(" Installing sentencepiece + numpy<2 (torch 2.1.2 needs numpy 1.x)...")
155
+ _r = subprocess.run(
156
+ [_py311, '-m', 'pip', 'install', '-q', 'sentencepiece', 'numpy<2'],
157
+ capture_output=True, text=True, timeout=120
158
+ )
159
+ if _r.returncode != 0:
160
+ print(f" sentencepiece/numpy install failed (non-fatal): {_r.stderr[-200:]}")
161
+
162
+ # Verify GPU works
163
+ _r = subprocess.run(
164
+ [_py311, '-c', 'import torch; t=torch.zeros(1,device="cuda"); print(f"torch {torch.__version__} CUDA {torch.version.cuda} — GPU OK")'],
165
+ capture_output=True, text=True, timeout=30
166
+ )
167
+ print(f" {_r.stdout.strip()}")
168
+ if 'OK' not in _r.stdout:
169
+ print(f" GPU test failed: {_r.stderr[-300:]}")
170
+ sys.exit(1)
171
+
172
+ # Re-exec under Python 3.11
173
+ print(f" Switching to Python 3.11...")
174
+ os.environ['_WYRM_P100_REEXEC'] = '1'
175
+ os.execv(_py311, [_py311] + sys.argv)
176
+ else:
177
+ print(f" GPU OK — {_gpu_name} is supported ✅")
178
+ except Exception as e:
179
+ print(f" GPU setup failed: {e}")
180
+ import traceback; traceback.print_exc()
181
+ sys.exit(1)
182
+ elif _P100_REEXEC:
183
+ print(f"Running under Python {sys.version.split()[0]} with P100-compatible PyTorch")
184
+
185
+ import torch
186
+
187
+ DUMMY_MODE_EARLY = '--dummy' in sys.argv or os.environ.get('WYRM_DUMMY') == '1'
188
+
189
+ if not torch.cuda.is_available():
190
+ if DUMMY_MODE_EARLY:
191
+ print("WARNING: No CUDA GPU. Running dummy mode on CPU.")
192
+ DEVICE = torch.device('cpu')
193
+ GPU_NAME = 'CPU'
194
+ GPU_VRAM = 0.0
195
+ IS_P100 = False
196
+ IS_T4 = False
197
+ DTYPE = torch.float32
198
+ else:
199
+ print("FATAL: No CUDA GPU detected. Exiting.")
200
+ print(" torch.cuda.is_available():", torch.cuda.is_available())
201
+ print(" This notebook requires a T4 GPU runtime.")
202
+ print(" (Use --dummy for CPU testing)")
203
+ sys.exit(1)
204
+ else:
205
+ DEVICE = torch.device('cuda')
206
+ GPU_NAME = torch.cuda.get_device_name(0)
207
+ GPU_VRAM = torch.cuda.get_device_properties(0).total_memory / (1024**3)
208
+ print(f"GPU: {GPU_NAME} ({GPU_VRAM:.1f} GB VRAM)")
209
+
210
+ # Test if GPU is actually usable
211
+ try:
212
+ _test = torch.zeros(1, device='cuda')
213
+ del _test
214
+ except RuntimeError as e:
215
+ print(f"FATAL: GPU detected but not usable: {e}")
216
+ print("This likely means the GPU's compute capability is not supported.")
217
+ print("Stop and restart for a compatible GPU.")
218
+ sys.exit(1)
219
+
220
+ IS_P100 = 'P100' in GPU_NAME or 'p100' in GPU_NAME.lower()
221
+ IS_T4 = 'T4' in GPU_NAME
222
+ print(f"GPU Type: {'T4' if IS_T4 else 'P100' if IS_P100 else 'Other'} — usable ✅")
223
+
224
+ # v26: FULL float32 — no autocast, no GradScaler, no corners
225
+ DTYPE = torch.float32
226
+ print(f"Precision: float32 (FULL — no mixed precision)")
227
+
228
+ # ═══════════════════════════════════════════════════════════════
229
+ # 1. ENVIRONMENT SETUP — copy from read-only input to working
230
+ # ═══════════════════════════════════════════════════════════════
231
+
232
+ IS_KAGGLE = os.path.exists('/kaggle')
233
+ DUMMY_MODE = '--dummy' in sys.argv or os.environ.get('WYRM_DUMMY') == '1'
234
+ LOCAL_TEST = '--local-test' in sys.argv or not IS_KAGGLE
235
+
236
+ if DUMMY_MODE:
237
+ print("\n*** DUMMY MODE — tiny model, ~60 second training ***\n")
238
+
239
+ if IS_KAGGLE:
240
+ INPUT_BASE = Path('/kaggle/input/wyrm-500m-base')
241
+ CKPT_INPUT = Path('/kaggle/input/wyrm-500m-ckpt')
242
+ WORKING = Path('/kaggle/working')
243
+ else:
244
+ # Local testing paths
245
+ _ws = Path(os.environ.get('WYRM_BASE', '/home/adam/workspace'))
246
+ INPUT_BASE = _ws / 'gladius_v2' / 'kaggle_upload'
247
+ CKPT_INPUT = _ws / 'gladius_v2' / '_ckpt_sim'
248
+ WORKING = Path(os.environ.get('WYRM_OUTPUT', '/tmp/wyrm_kaggle_test'))
249
+
250
+ OUTPUT = WORKING / 'output'
251
+ RUNS = OUTPUT / 'runs' / 'wyrm-500m'
252
+ CKPT_OUT = OUTPUT / 'checkpoint'
253
+ TELEM = OUTPUT / 'telemetry'
254
+
255
+ for d in [OUTPUT, RUNS, CKPT_OUT, TELEM]:
256
+ d.mkdir(parents=True, exist_ok=True)
257
+
258
+ # ── Copy kernel + staging from input to working (READ-ONLY FIX) ──
259
+ KERNEL_DST = WORKING / 'kernel'
260
+ STAGING_DST = WORKING / 'staging'
261
+
262
+ def copy_tree_safe(src: Path, dst: Path, label: str):
263
+ """Copy directory tree, overwriting if exists."""
264
+ if not src.exists():
265
+ print(f" WARNING: {label} source not found: {src}")
266
+ return False
267
+ if dst.exists():
268
+ shutil.rmtree(dst)
269
+ shutil.copytree(src, dst)
270
+ n_files = sum(1 for _ in dst.rglob('*') if _.is_file())
271
+ print(f" Copied {label}: {src} -> {dst} ({n_files} files)")
272
+ return True
273
+
274
+ print("\nCopying code to writable location...")
275
+
276
+ # Debug: list what's at input path
277
+ if IS_KAGGLE:
278
+ import subprocess
279
+ print(f" Input base: {INPUT_BASE}")
280
+ print(f" /kaggle/input/ contents:")
281
+ kaggle_input = Path('/kaggle/input')
282
+ if kaggle_input.exists():
283
+ for p in sorted(kaggle_input.iterdir()):
284
+ print(f" {p.name}: {'DIR' if p.is_dir() else 'FILE'}")
285
+ if p.is_dir():
286
+ for pp in sorted(p.iterdir()):
287
+ print(f" {pp.name}: {'DIR' if pp.is_dir() else 'FILE'}")
288
+ if pp.is_dir():
289
+ for ppp in sorted(pp.iterdir())[:5]:
290
+ print(f" {ppp.name}: {'DIR' if ppp.is_dir() else 'FILE'}")
291
+ else:
292
+ print(f" /kaggle/input does not exist!")
293
+
294
+ if INPUT_BASE.exists():
295
+ print(f" Dataset contents:")
296
+ for p in sorted(INPUT_BASE.rglob('*')):
297
+ if p.is_file():
298
+ rel = p.relative_to(INPUT_BASE)
299
+ print(f" {rel}: {p.stat().st_size}B")
300
+ else:
301
+ print(f" WARNING: {INPUT_BASE} does not exist!")
302
+ # Search for kernel package anywhere under /kaggle/input
303
+ for candidate in Path('/kaggle/input').rglob('kernel/__init__.py'):
304
+ found_base = candidate.parent.parent
305
+ print(f" FOUND kernel at: {found_base}")
306
+ INPUT_BASE = found_base
307
+ break
308
+
309
+ copy_tree_safe(INPUT_BASE / 'kernel', KERNEL_DST, 'kernel')
310
+ copy_tree_safe(INPUT_BASE / 'staging', STAGING_DST, 'staging')
311
+
312
+ # ── Copy tokenizer if present ──
313
+ _tok_src = INPUT_BASE / 'gladius_bpe_16k.model'
314
+ _tok_dst = WORKING / 'gladius_bpe_16k.model'
315
+ if _tok_src.exists():
316
+ shutil.copy2(_tok_src, _tok_dst)
317
+ print(f" Tokenizer copied: {_tok_dst} ({_tok_src.stat().st_size}B)")
318
+ else:
319
+ # Search recursively under INPUT_BASE
320
+ for _tk in INPUT_BASE.rglob('gladius_bpe_16k.model'):
321
+ shutil.copy2(_tk, _tok_dst)
322
+ print(f" Tokenizer found at {_tk}, copied to {_tok_dst}")
323
+ break
324
+ else:
325
+ print(" WARNING: gladius_bpe_16k.model not found in dataset — BPE will use procedural only")
326
+
327
+ # ── Setup import paths (writable copies only) ──
328
+ sys.path.insert(0, str(WORKING))
329
+ sys.path.insert(0, str(STAGING_DST))
330
+
331
+ # ── Verify imports work ──
332
+ try:
333
+ from kernel.config import KernelConfig
334
+ from kernel.kernel import GladiusKernel
335
+ from kernel.gaussian_head import (
336
+ GaussianConfig, GaussianSpecialist, GaussianVQVAE,
337
+ ProceduralGaussianDataset, render_gaussians_2d
338
+ )
339
+ print(" Kernel imports: OK")
340
+ except ImportError as e:
341
+ print(f"FATAL: Kernel import failed: {e}")
342
+ traceback.print_exc()
343
+ sys.exit(1)
344
+
345
+ # ══════════════════════��════════════════════════════════════════
346
+ # 2. CONFIGURATION
347
+ # ═══════════════════════════════════════════════════════════════
348
+
349
+ import torch.nn as nn
350
+ import torch.nn.functional as F
351
+
352
+
353
+ @dataclasses.dataclass
354
+ class WyrmKaggleConfig:
355
+ """Kaggle-optimized WYRM training configuration."""
356
+
357
+ # ── Architecture ──
358
+ hidden_dim: int = 1024
359
+ num_layers: int = 24
360
+ num_heads: int = 32
361
+ head_dim: int = 32
362
+ ffn_dim: int = 4096
363
+ max_seq_len: int = 1024
364
+ vocab_size: int = 32000 # BPE
365
+
366
+ # ── Memory / specialist ──
367
+ hot_memory_slots: int = 1024
368
+ warm_rank: int = 64
369
+ time_dim: int = 128
370
+ time_num_frequencies: int = 16
371
+ time_max_events: int = 64
372
+ cognition_state_dim: int = 256
373
+ cognition_modes: int = 4
374
+ cognition_prompt_types: int = 5
375
+ register_dim: int = 4
376
+ intent_dim: int = 4
377
+ max_tools: int = 64
378
+ num_specialists: int = 5 # reasoning, math, code, general, gaussian
379
+ router_top_k: int = 2
380
+ attention_sparse_budget: int = 128
381
+
382
+ # ── Training ──
383
+ batch_size: int = 4 # v27: 4 batch (5.4/15.9 GB VRAM = headroom)
384
+ grad_accumulation: int = 4 # v27: 4 * 4 = 16 effective batch
385
+ total_steps: int = 15000 # Full training run (chain across sessions)
386
+ warmup_steps: int = 200
387
+ max_grad_norm: float = 0.5
388
+ weight_decay: float = 0.01
389
+ seed: int = 42
390
+
391
+ # ── Learning rates ──
392
+ lr_backbone: float = 3e-5
393
+ lr_specialists: float = 3e-4
394
+ lr_router: float = 1e-3
395
+ lr_tools: float = 5e-4
396
+ lr_depth: float = 3e-4
397
+ lr_embeddings: float = 3e-4
398
+
399
+ # ── Loss weights ──
400
+ balance_loss_weight: float = 0.05
401
+ cognition_loss_weight: float = 0.1
402
+ aux_loss_weight: float = 0.3
403
+
404
+ # ── Curriculum phases ──
405
+ phase1_end: int = 5000 # Foundation
406
+ phase2_end: int = 10000 # Reasoning
407
+ phase3_end: int = 13000 # Depth
408
+ # phase4: Omega (remainder)
409
+
410
+ # ── Depth ratios: [D1, D2, D3, D4, D5] ──
411
+ foundation_depth: tuple = (0.40, 0.40, 0.20, 0.00, 0.00)
412
+ reasoning_depth: tuple = (0.15, 0.15, 0.25, 0.25, 0.20)
413
+ depth_depth: tuple = (0.05, 0.05, 0.20, 0.20, 0.50)
414
+ omega_depth: tuple = (0.20, 0.20, 0.20, 0.20, 0.20)
415
+
416
+ # ── Domain ratios ──
417
+ bonus_ratio: float = 0.15
418
+ language_ratio: float = 0.40
419
+ science_ratio: float = 0.05
420
+ gaussian_ratio: float = 0.02
421
+
422
+ # ── Checkpointing ──
423
+ checkpoint_every: int = 250 # More frequent saves for 12h session chaining
424
+ log_every: int = 10 # Every 10 steps
425
+ eval_every: int = 250
426
+
427
+ # ── NaN handling ──
428
+ max_nan_retries: int = 5
429
+ max_total_nans: int = 50
430
+
431
+ # ── VRAM safety ──
432
+ vram_warning_gb: float = 13.0
433
+ min_disk_gb: float = 3.0
434
+
435
+ def get_phase(self, step: int) -> str:
436
+ if step < self.phase1_end: return 'foundation'
437
+ if step < self.phase2_end: return 'reasoning'
438
+ if step < self.phase3_end: return 'depth'
439
+ return 'omega'
440
+
441
+ def get_depth_ratios(self, step: int) -> tuple:
442
+ phase = self.get_phase(step)
443
+ return getattr(self, f'{phase}_depth')
444
+
445
+ def get_cognitive_ratio(self, step: int) -> float:
446
+ return max(0.0, 1.0 - self.bonus_ratio - self.language_ratio
447
+ - self.science_ratio - self.gaussian_ratio)
448
+
449
+
450
+ def make_dummy_config() -> WyrmKaggleConfig:
451
+ """Tiny config for testing — same code path, finishes in ~60s."""
452
+ return WyrmKaggleConfig(
453
+ hidden_dim=64,
454
+ num_layers=2,
455
+ num_heads=4, # min 4 for Synthase depth KV heads
456
+ head_dim=16, # 64 / 4 = 16
457
+ ffn_dim=256,
458
+ max_seq_len=128,
459
+ vocab_size=1000,
460
+ hot_memory_slots=32,
461
+ warm_rank=4,
462
+ time_dim=32,
463
+ time_num_frequencies=8,
464
+ time_max_events=16,
465
+ cognition_state_dim=32,
466
+ max_tools=8,
467
+ num_specialists=5,
468
+ router_top_k=2,
469
+ attention_sparse_budget=32,
470
+ batch_size=2,
471
+ grad_accumulation=2,
472
+ total_steps=50,
473
+ warmup_steps=5,
474
+ checkpoint_every=25,
475
+ log_every=1,
476
+ eval_every=25,
477
+ phase1_end=15,
478
+ phase2_end=30,
479
+ phase3_end=40,
480
+ )
481
+
482
+
483
+ # ═══════════════════════════════════════════════════════════════
484
+ # 3. TOKENIZERS (inline — no external dependency)
485
+ # ═══════════════════════════════════════════════════════════════
486
+
487
+ class MathTokenizer:
488
+ """Structural tokenizer for mathematical expressions. 128 tokens."""
489
+ VOCAB_SIZE = 128
490
+ PAD_ID = 0; BOS_ID = 1; EOS_ID = 2
491
+
492
+ def __init__(self):
493
+ self.token_to_id = {}; self.id_to_token = {}
494
+ idx = 0
495
+ for name in ['[PAD_MATH]', '[BOS_MATH]', '[EOS_MATH]']:
496
+ self.token_to_id[name] = idx; self.id_to_token[idx] = name; idx += 1
497
+ for i in range(10):
498
+ s = str(i); self.token_to_id[s] = idx; self.id_to_token[idx] = s; idx += 1
499
+ for op in ['+','-','*','/','=','^','!','%','<','>','.',':','|']:
500
+ self.token_to_id[op] = idx; self.id_to_token[idx] = op; idx += 1
501
+ for i in range(26):
502
+ ch = chr(ord('a')+i); self.token_to_id[ch] = idx; self.id_to_token[idx] = ch; idx += 1
503
+ for s in ['(',')',',','_',' ','\n','STEP','OUT','FILL','NEXT','GIVEN','QED']:
504
+ self.token_to_id[s] = idx; self.id_to_token[idx] = s; idx += 1
505
+
506
+ @property
507
+ def vocab_size(self): return self.VOCAB_SIZE
508
+
509
+ def encode(self, text: str, add_special: bool = True) -> List[int]:
510
+ tokens = [self.BOS_ID] if add_special else []
511
+ for ch in text.strip():
512
+ if ch in self.token_to_id: tokens.append(self.token_to_id[ch])
513
+ if add_special: tokens.append(self.EOS_ID)
514
+ return tokens
515
+
516
+ def decode(self, ids: List[int]) -> str:
517
+ return ''.join(self.id_to_token.get(i, '?') for i in ids if i > 2)
518
+
519
+
520
+ class ByteTokenizer:
521
+ """Byte-level tokenizer: 256 bytes + 3 specials = 259."""
522
+ PAD_ID = 256; BOS_ID = 257; EOS_ID = 258; VOCAB_SIZE = 259
523
+
524
+ def encode(self, data: bytes) -> List[int]:
525
+ return [self.BOS_ID] + [int(b) for b in data] + [self.EOS_ID]
526
+
527
+ def decode(self, ids: List[int]) -> bytes:
528
+ return bytes(i for i in ids if 0 <= i <= 255)
529
+
530
+
531
+ # ═══════════════════════════════════════════════════════════════
532
+ # 4. MULTI-TOKENIZER EMBEDDING
533
+ # ═══════════════════════════════════════════════════════════════
534
+
535
+ class MultiTokenizerEmbedding(nn.Module):
536
+ """Three embedding tables: BPE (32K), Math (128), Byte (259)."""
537
+
538
+ def __init__(self, hidden_dim: int, bpe_vocab_size: int = 32000):
539
+ super().__init__()
540
+ self.hidden_dim = hidden_dim
541
+ self.bpe_vocab_size = bpe_vocab_size
542
+ self.math_vocab_size = 128
543
+ self.byte_vocab_size = 259
544
+
545
+ self.bpe_embed = nn.Embedding(bpe_vocab_size, hidden_dim, padding_idx=0)
546
+ self.math_embed = nn.Embedding(128, hidden_dim, padding_idx=0)
547
+ self.byte_embed = nn.Embedding(259, hidden_dim, padding_idx=256)
548
+
549
+ self.bpe_head = nn.Linear(hidden_dim, bpe_vocab_size, bias=False)
550
+ self.math_head = nn.Linear(hidden_dim, 128, bias=False)
551
+ self.byte_head = nn.Linear(hidden_dim, 259, bias=False)
552
+
553
+ self.bpe_head.weight = self.bpe_embed.weight # Weight tying
554
+ self.scale = math.sqrt(hidden_dim)
555
+
556
+ self._init_weights()
557
+
558
+ def _init_weights(self):
559
+ for e in [self.bpe_embed, self.math_embed, self.byte_embed]:
560
+ nn.init.normal_(e.weight, std=0.02)
561
+ with torch.no_grad():
562
+ self.bpe_embed.weight[0].zero_()
563
+ self.math_embed.weight[0].zero_()
564
+ self.byte_embed.weight[256].zero_()
565
+ nn.init.xavier_uniform_(self.math_head.weight)
566
+ nn.init.xavier_uniform_(self.byte_head.weight)
567
+
568
+ def embed(self, ids: torch.Tensor, domain: str) -> torch.Tensor:
569
+ if domain == 'bpe': return self.bpe_embed(ids) * self.scale
570
+ elif domain == 'math': return self.math_embed(ids) * self.scale
571
+ elif domain == 'byte': return self.byte_embed(ids) * self.scale
572
+ raise ValueError(f"Unknown domain: {domain}")
573
+
574
+ def project(self, hidden: torch.Tensor, domain: str) -> torch.Tensor:
575
+ if domain == 'bpe': return self.bpe_head(hidden)
576
+ elif domain == 'math': return self.math_head(hidden)
577
+ elif domain == 'byte': return self.byte_head(hidden)
578
+ raise ValueError(f"Unknown domain: {domain}")
579
+
580
+ def get_vocab_size(self, domain: str) -> int:
581
+ return {'bpe': self.bpe_vocab_size, 'math': 128, 'byte': 259}[domain]
582
+
583
+
584
+ # ═══════════════════════════════════════════════════════════════
585
+ # 5. DOMAIN MEMBRANES (Plug)
586
+ # ═══════════════════════════════════════════════════════════════
587
+
588
+ class DomainMembrane(nn.Module):
589
+ def __init__(self, hidden_dim):
590
+ super().__init__()
591
+ self.proj = nn.Linear(hidden_dim, hidden_dim)
592
+ self.norm = nn.LayerNorm(hidden_dim)
593
+ nn.init.eye_(self.proj.weight)
594
+ with torch.no_grad(): self.proj.weight.add_(torch.randn_like(self.proj.weight)*0.01)
595
+ nn.init.zeros_(self.proj.bias)
596
+
597
+ def forward(self, x): return self.norm(self.proj(x))
598
+
599
+
600
+ class PlugMembranes(nn.Module):
601
+ def __init__(self, hidden_dim):
602
+ super().__init__()
603
+ self.bpe_membrane = DomainMembrane(hidden_dim)
604
+ self.math_membrane = DomainMembrane(hidden_dim)
605
+ self.byte_membrane = DomainMembrane(hidden_dim)
606
+
607
+ def forward(self, x, domain):
608
+ return getattr(self, f'{domain}_membrane')(x)
609
+
610
+
611
+ # ═══════════════════════════════════════════════════════════════
612
+ # 6. AUXILIARY PREDICTION HEAD (L7)
613
+ # ═══════════════════════════════════════════════════════════════
614
+
615
+ class AuxiliaryPredictionHead(nn.Module):
616
+ def __init__(self, hidden_dim: int, vocab_size: int):
617
+ super().__init__()
618
+ self.norm = nn.LayerNorm(hidden_dim)
619
+ self.proj = nn.Linear(hidden_dim, vocab_size, bias=False)
620
+ nn.init.xavier_uniform_(self.proj.weight)
621
+
622
+ def forward(self, hidden): return self.proj(self.norm(hidden))
623
+
624
+
625
+ # ═══════════════════════════════════════════════════════════════
626
+ # 7. DATA — Procedural curriculum (no external files for dummy)
627
+ # ═══════════════════════════════════════════════════════════════
628
+
629
+ def load_bpe_tokenizer(path: str):
630
+ """Load SentencePiece BPE tokenizer."""
631
+ try:
632
+ import sentencepiece as spm
633
+ sp = spm.SentencePieceProcessor()
634
+ sp.load(path)
635
+ print(f" BPE tokenizer loaded: {path} (vocab={sp.get_piece_size()})")
636
+ return sp
637
+ except Exception as e:
638
+ print(f" WARNING: BPE tokenizer failed: {e}")
639
+ return None
640
+
641
+
642
+ class ProceduralTextDataset:
643
+ """Generates synthetic text data for training. No file dependencies."""
644
+
645
+ TEMPLATES = [
646
+ "The {adj} {noun} {verb} the {adj2} {noun2}.",
647
+ "In the {noun}, a {adj} {noun2} was {verb_ing}.",
648
+ "{noun} and {noun2} are fundamentally {adj}.",
649
+ "The concept of {noun} relates to {noun2} through {adj} mechanisms.",
650
+ "When {noun} encounters {noun2}, the result is {adj}.",
651
+ ]
652
+ ADJS = ['quick','large','small','bright','dark','complex','simple','ancient','modern','curious']
653
+ NOUNS = ['system','pattern','network','structure','function','equation','theory','model','field','space']
654
+ VERBS = ['transforms','analyzes','computes','observes','generates','processes','evaluates','maps']
655
+ VERB_ING = ['computing','analyzing','transforming','processing','generating','mapping','observing']
656
+
657
+ def __init__(self, tokenizer, seq_len: int):
658
+ self.tok = tokenizer
659
+ self.seq_len = seq_len
660
+
661
+ def sample(self) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
662
+ """Generate a random text sample, tokenize, return (input, target, mask)."""
663
+ # Generate enough text to fill seq_len
664
+ text_parts = []
665
+ while True:
666
+ tmpl = random.choice(self.TEMPLATES)
667
+ text = tmpl.format(
668
+ adj=random.choice(self.ADJS), adj2=random.choice(self.ADJS),
669
+ noun=random.choice(self.NOUNS), noun2=random.choice(self.NOUNS),
670
+ verb=random.choice(self.VERBS), verb_ing=random.choice(self.VERB_ING),
671
+ )
672
+ text_parts.append(text)
673
+ full_text = ' '.join(text_parts)
674
+ if self.tok is not None:
675
+ ids = self.tok.encode(full_text)
676
+ else:
677
+ ids = [ord(c) % 1000 for c in full_text] # Fallback for dummy
678
+ if len(ids) >= self.seq_len + 1:
679
+ break
680
+
681
+ ids = ids[:self.seq_len + 1]
682
+ x = torch.tensor(ids[:-1], dtype=torch.long)
683
+ y = torch.tensor(ids[1:], dtype=torch.long)
684
+ m = torch.ones(len(x), dtype=torch.float32)
685
+
686
+ # Pad if needed
687
+ if len(x) < self.seq_len:
688
+ pad = self.seq_len - len(x)
689
+ x = F.pad(x, (0, pad))
690
+ y = F.pad(y, (0, pad))
691
+ m = F.pad(m, (0, pad))
692
+
693
+ m[y == 0] = 0.0
694
+ return x[:self.seq_len], y[:self.seq_len], m[:self.seq_len]
695
+
696
+
697
+ class ProceduralMathDataset:
698
+ """Generates synthetic math problems."""
699
+
700
+ def __init__(self, seq_len: int):
701
+ self.tok = MathTokenizer()
702
+ self.seq_len = seq_len
703
+
704
+ def sample(self) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
705
+ # Generate random arithmetic
706
+ a, b = random.randint(1, 100), random.randint(1, 100)
707
+ op = random.choice(['+', '-', '*'])
708
+ result = eval(f"{a}{op}{b}")
709
+ expr = f"D2|GIVEN:{a}{op}{b}|OUT:{result}"
710
+ ids = self.tok.encode(expr)
711
+ if len(ids) > self.seq_len + 1:
712
+ ids = ids[:self.seq_len + 1]
713
+ x = ids[:-1]; y = ids[1:]
714
+ pad = self.seq_len - len(x)
715
+ mask = [1.0]*len(x) + [0.0]*pad
716
+ x = x + [0]*pad; y = y + [0]*pad
717
+ for i in range(len(mask)):
718
+ if y[i] == 0: mask[i] = 0.0
719
+ return (torch.tensor(x[:self.seq_len], dtype=torch.long),
720
+ torch.tensor(y[:self.seq_len], dtype=torch.long),
721
+ torch.tensor(mask[:self.seq_len], dtype=torch.float32))
722
+
723
+
724
+ class UnifiedSampler:
725
+ """Samples from curriculum based on phase ratios."""
726
+
727
+ def __init__(self, cfg: WyrmKaggleConfig, text_ds, math_ds, gaussian_ds=None):
728
+ self.cfg = cfg
729
+ self.text_ds = text_ds
730
+ self.math_ds = math_ds
731
+ self.gaussian_ds = gaussian_ds
732
+
733
+ def _pick_difficulty(self, step: int) -> int:
734
+ ratios = self.cfg.get_depth_ratios(step)
735
+ r = random.random()
736
+ cum = 0.0
737
+ for i, ratio in enumerate(ratios):
738
+ cum += ratio
739
+ if r < cum: return i + 1
740
+ return 5
741
+
742
+ def sample_batch(self, batch_size: int, step: int) -> Dict[str, Any]:
743
+ """Sample a batch based on curriculum ratios."""
744
+ roll = random.random()
745
+ cog_ratio = self.cfg.get_cognitive_ratio(step)
746
+
747
+ if roll < cog_ratio:
748
+ # Cognitive / math
749
+ difficulty = self._pick_difficulty(step)
750
+ if random.random() < 0.3 and self.math_ds is not None:
751
+ samples = [self.math_ds.sample() for _ in range(batch_size)]
752
+ domain = 'math'
753
+ task = f'math_D{difficulty}'
754
+ else:
755
+ samples = [self.text_ds.sample() for _ in range(batch_size)]
756
+ domain = 'bpe'
757
+ task = f'cog_D{difficulty}'
758
+ elif roll < cog_ratio + self.cfg.gaussian_ratio and self.gaussian_ds is not None:
759
+ # Gaussian
760
+ samples = [self.text_ds.sample() for _ in range(batch_size)]
761
+ domain = 'bpe'
762
+ task = 'gaussian'
763
+ else:
764
+ # Language / science / bonus
765
+ samples = [self.text_ds.sample() for _ in range(batch_size)]
766
+ domain = 'bpe'
767
+ task = 'language'
768
+
769
+ return {
770
+ 'input_ids': torch.stack([s[0] for s in samples]),
771
+ 'target_ids': torch.stack([s[1] for s in samples]),
772
+ 'loss_mask': torch.stack([s[2] for s in samples]),
773
+ 'task_type': task,
774
+ 'domain': domain,
775
+ 'difficulty': self._pick_difficulty(step),
776
+ 'corpus': task,
777
+ }
778
+
779
+
780
+ # ═══════════════════════════════════════════════════════════════
781
+ # 8. COSINE WARMUP SCHEDULER
782
+ # ═══════════════════════════════════════════════════════════════
783
+
784
+ class CosineWarmupScheduler:
785
+ def __init__(self, optimizer, warmup_steps: int, total_steps: int):
786
+ self.optimizer = optimizer
787
+ self.warmup_steps = warmup_steps
788
+ self.total_steps = total_steps
789
+ self.base_lrs = [pg['lr'] for pg in optimizer.param_groups]
790
+
791
+ def step(self, current_step: int):
792
+ if current_step < self.warmup_steps:
793
+ factor = current_step / max(1, self.warmup_steps)
794
+ else:
795
+ progress = (current_step - self.warmup_steps) / max(1, self.total_steps - self.warmup_steps)
796
+ factor = 0.1 + 0.9 * 0.5 * (1.0 + math.cos(math.pi * progress))
797
+ for pg, base_lr in zip(self.optimizer.param_groups, self.base_lrs):
798
+ pg['lr'] = base_lr * factor
799
+
800
+
801
+ # ═══════════════════════════════════════════════════════════════
802
+ # 9. KAGGLE SESSION MANAGEMENT
803
+ # ═══════════════════════════════════════════════════════════════
804
+
805
+ class KaggleSessionManager:
806
+ """Manages Kaggle session time limits and checkpoint chaining."""
807
+
808
+ MAX_RUNTIME_S = 9 * 3600 # 9 hours (Kaggle limit ~12h, save buffer)
809
+ SAVE_BUFFER_S = 600 # Start saving 10 min before cutoff
810
+
811
+ def __init__(self, output_dir: Path, ckpt_input: Path):
812
+ self.start_time = time.time()
813
+ self.output_dir = output_dir
814
+ self.ckpt_input = ckpt_input
815
+ self.chain_pos = 0
816
+ self.resume_path = None
817
+ self._detect_chain()
818
+
819
+ def _detect_chain(self):
820
+ """Find latest checkpoint from previous session."""
821
+ for ckpt_dir in [self.ckpt_input, CKPT_OUT]:
822
+ if not ckpt_dir.exists():
823
+ continue
824
+ manifest = ckpt_dir / 'manifest.json'
825
+ if manifest.exists():
826
+ try:
827
+ info = json.loads(manifest.read_text())
828
+ self.chain_pos = info.get('chain', 0) + 1
829
+ ckpt_file = ckpt_dir / info.get('file', 'latest.pt')
830
+ if ckpt_file.exists():
831
+ self.resume_path = str(ckpt_file)
832
+ print(f" Chain #{self.chain_pos}: resuming from step {info.get('step', '?')}")
833
+ return
834
+ except: pass
835
+ for name in ['latest.pt', 'checkpoint_latest.pt']:
836
+ p = ckpt_dir / name
837
+ if p.exists():
838
+ self.resume_path = str(p)
839
+ print(f" Found checkpoint: {p}")
840
+ return
841
+
842
+ def should_save_and_exit(self) -> bool:
843
+ elapsed = time.time() - self.start_time
844
+ return elapsed > (self.MAX_RUNTIME_S - self.SAVE_BUFFER_S)
845
+
846
+ def elapsed_hours(self) -> float:
847
+ return (time.time() - self.start_time) / 3600
848
+
849
+ def save_manifest(self, step: int, filename: str):
850
+ manifest = {
851
+ 'chain': self.chain_pos,
852
+ 'step': step,
853
+ 'file': filename,
854
+ 'timestamp': datetime.datetime.now().isoformat(),
855
+ 'elapsed_hours': self.elapsed_hours(),
856
+ }
857
+ (CKPT_OUT / 'manifest.json').write_text(json.dumps(manifest, indent=2))
858
+ print(f" Manifest saved: chain={self.chain_pos}, step={step}")
859
+
860
+
861
+ # ═══════════════════════════════════════════════════════════════
862
+ # 10. VRAM MONITORING
863
+ # ═══════════════════════════════════════════════════════════════
864
+
865
+ def vram_gb() -> float:
866
+ if torch.cuda.is_available():
867
+ return torch.cuda.memory_allocated() / (1024**3)
868
+ return 0.0
869
+
870
+ def vram_total_gb() -> float:
871
+ if torch.cuda.is_available():
872
+ return torch.cuda.get_device_properties(0).total_memory / (1024**3)
873
+ return 0.0
874
+
875
+ def vram_cleanup():
876
+ """Emergency VRAM cleanup."""
877
+ if torch.cuda.is_available():
878
+ torch.cuda.empty_cache()
879
+ torch.cuda.synchronize()
880
+
881
+ def disk_free_gb(path: str = '/') -> float:
882
+ st = os.statvfs(path)
883
+ return (st.f_bavail * st.f_frsize) / (1024**3)
884
+
885
+
886
+ # ═══════════════════════════════════════════════════════════════
887
+ # 11. TRAINER
888
+ # ═══════════════════════════════════════════════════════════════
889
+
890
+ class WyrmKaggleTrainer:
891
+ """Complete WYRM training loop for Kaggle."""
892
+
893
+ def __init__(self, cfg: WyrmKaggleConfig, device: torch.device):
894
+ self.cfg = cfg
895
+ self.device = device
896
+ # v26: full fp32 — no dtype switching
897
+
898
+ # ── Build kernel ──
899
+ self.kernel_config = KernelConfig(
900
+ vocab_size=cfg.vocab_size,
901
+ hidden_dim=cfg.hidden_dim,
902
+ num_layers=cfg.num_layers,
903
+ num_heads=cfg.num_heads,
904
+ head_dim=cfg.head_dim,
905
+ ffn_dim=cfg.ffn_dim,
906
+ max_seq_len=cfg.max_seq_len,
907
+ hot_memory_slots=cfg.hot_memory_slots,
908
+ warm_rank=cfg.warm_rank,
909
+ cold_embedding_dim=cfg.hidden_dim,
910
+ time_dim=cfg.time_dim,
911
+ time_num_frequencies=cfg.time_num_frequencies,
912
+ time_max_events=cfg.time_max_events,
913
+ cognition_state_dim=cfg.cognition_state_dim,
914
+ cognition_modes=cfg.cognition_modes,
915
+ cognition_prompt_types=cfg.cognition_prompt_types,
916
+ register_dim=cfg.register_dim,
917
+ intent_dim=cfg.intent_dim,
918
+ max_tools=cfg.max_tools,
919
+ num_specialists=cfg.num_specialists,
920
+ router_top_k=cfg.router_top_k,
921
+ attention_sparse_budget=cfg.attention_sparse_budget,
922
+ batch_size=cfg.batch_size,
923
+ accumulation_steps=cfg.grad_accumulation,
924
+ seed=cfg.seed,
925
+ )
926
+
927
+ print(f"\nBuilding WYRM kernel...")
928
+ self.model = GladiusKernel(self.kernel_config).to(device)
929
+
930
+ # ── Multi-tokenizer embedding ──
931
+ self.multi_embed = MultiTokenizerEmbedding(cfg.hidden_dim, cfg.vocab_size).to(device)
932
+ with torch.no_grad():
933
+ self.multi_embed.bpe_embed.weight.copy_(self.model.embeddings.token_embed.weight)
934
+ self.model.embeddings.token_embed = self.multi_embed.bpe_embed
935
+ self.model.embeddings.output_head = self.multi_embed.bpe_head
936
+ self.model.embeddings.scale = self.multi_embed.scale
937
+
938
+ # ── Plug membranes ──
939
+ self.plug_membranes = PlugMembranes(cfg.hidden_dim).to(device)
940
+
941
+ # ── L7 Aux head ──
942
+ self.aux_head = AuxiliaryPredictionHead(cfg.hidden_dim, cfg.vocab_size).to(device)
943
+
944
+ # ── Gaussian specialist head ──
945
+ self.gaussian_config = GaussianConfig(
946
+ num_anchors=16 if DUMMY_MODE else 64,
947
+ details_per_anchor=8 if DUMMY_MODE else 32,
948
+ codebook_size=512 if DUMMY_MODE else 4096,
949
+ codebook_dim=32 if DUMMY_MODE else 64,
950
+ )
951
+ self.gaussian_head = GaussianSpecialist(
952
+ backbone_dim=cfg.hidden_dim,
953
+ num_backbone_layers=cfg.num_layers,
954
+ config=self.gaussian_config,
955
+ ).to(device)
956
+
957
+ # ── Gradient checkpointing (saves ~40% VRAM) ──
958
+ if not DUMMY_MODE:
959
+ from torch.utils.checkpoint import checkpoint as torch_checkpoint
960
+ for layer in self.model.layers:
961
+ orig = layer.forward
962
+ def make_ckpt(fwd):
963
+ def wrapped(*a, **kw):
964
+ return torch_checkpoint(fwd, *a, use_reentrant=False, **kw)
965
+ return wrapped
966
+ layer.forward = make_ckpt(orig)
967
+
968
+ # ── L7 hook ──
969
+ self._layer7_output = None
970
+ def _capture_l7(module, inp, out):
971
+ self._layer7_output = out[0] if isinstance(out, tuple) else out
972
+ if len(self.model.layers) > 7:
973
+ self.model.layers[7].register_forward_hook(_capture_l7)
974
+ elif len(self.model.layers) > 0:
975
+ self.model.layers[-1].register_forward_hook(_capture_l7)
976
+
977
+ # ── v26: No GradScaler — full fp32 ──
978
+ self.scaler = None
979
+
980
+ # ── Report params ──
981
+ total = sum(p.numel() for p in self.model.parameters())
982
+ embed_p = sum(p.numel() for p in self.multi_embed.parameters())
983
+ plug_p = sum(p.numel() for p in self.plug_membranes.parameters())
984
+ aux_p = sum(p.numel() for p in self.aux_head.parameters())
985
+ gauss_p = sum(p.numel() for p in self.gaussian_head.parameters())
986
+ grand_total = total + embed_p + plug_p + aux_p + gauss_p
987
+ print(f" Kernel: {total:,} | MultiEmbed: {embed_p:,} | Plug: {plug_p:,}")
988
+ print(f" AuxHead: {aux_p:,} | GaussianHead: {gauss_p:,}")
989
+ print(f" TOTAL: {grand_total:,} params")
990
+ print(f" VRAM after model: {vram_gb():.2f} GB / {vram_total_gb():.1f} GB")
991
+
992
+ # ── State ──
993
+ self.step = 0
994
+ self.best_loss = float('inf')
995
+ self.nan_count = 0
996
+ self.total_nan_count = 0
997
+ self.telemetry = []
998
+
999
+ def setup_optimizer(self):
1000
+ """Differential LR optimizer."""
1001
+ groups = defaultdict(list)
1002
+
1003
+ for name, param in self.model.named_parameters():
1004
+ if not param.requires_grad: continue
1005
+ if 'depth' in name or 'synthase' in name: groups['depth'].append(param)
1006
+ elif 'specialist' in name: groups['specialists'].append(param)
1007
+ elif 'router' in name: groups['router'].append(param)
1008
+ elif 'tool_cortex' in name: groups['tools'].append(param)
1009
+ elif 'embed' in name or 'output_head' in name: groups['embed_bpe'].append(param)
1010
+ else: groups['backbone'].append(param)
1011
+
1012
+ for name, param in self.multi_embed.named_parameters():
1013
+ if not param.requires_grad: continue
1014
+ if 'math' in name: groups['embed_math'].append(param)
1015
+ elif 'byte' in name: groups['embed_byte'].append(param)
1016
+
1017
+ lr_map = {
1018
+ 'backbone': self.cfg.lr_backbone,
1019
+ 'specialists': self.cfg.lr_specialists,
1020
+ 'router': self.cfg.lr_router,
1021
+ 'tools': self.cfg.lr_tools,
1022
+ 'depth': self.cfg.lr_depth,
1023
+ 'embed_bpe': self.cfg.lr_embeddings,
1024
+ 'embed_math': self.cfg.lr_embeddings * 3,
1025
+ 'embed_byte': self.cfg.lr_embeddings * 3,
1026
+ }
1027
+
1028
+ param_groups = []
1029
+ for name, params in groups.items():
1030
+ if params:
1031
+ param_groups.append({'params': params, 'lr': lr_map.get(name, 3e-5), 'name': name})
1032
+
1033
+ # Aux head, Plug membranes, Gaussian head
1034
+ param_groups.append({'params': list(self.aux_head.parameters()),
1035
+ 'lr': self.cfg.lr_backbone * 3, 'name': 'aux_head'})
1036
+ param_groups.append({'params': list(self.plug_membranes.parameters()),
1037
+ 'lr': self.cfg.lr_specialists, 'name': 'plug_membranes'})
1038
+ param_groups.append({'params': list(self.gaussian_head.parameters()),
1039
+ 'lr': self.cfg.lr_specialists, 'name': 'gaussian_head'})
1040
+
1041
+ self.optimizer = torch.optim.AdamW(
1042
+ param_groups, weight_decay=self.cfg.weight_decay, betas=(0.9, 0.95))
1043
+ self.scheduler = CosineWarmupScheduler(
1044
+ self.optimizer, self.cfg.warmup_steps, self.cfg.total_steps)
1045
+
1046
+ n_groups = len([g for g in param_groups if g['params']])
1047
+ n_params = sum(sum(p.numel() for p in g['params']) for g in param_groups)
1048
+ print(f" Optimizer: {n_groups} groups, {n_params:,} params")
1049
+
1050
+ def compute_loss(self, batch: Dict[str, Any]) -> Tuple[torch.Tensor, Dict[str, float]]:
1051
+ """Forward pass + loss computation."""
1052
+ input_ids = batch['input_ids'].to(self.device)
1053
+ target_ids = batch['target_ids'].to(self.device)
1054
+ loss_mask = batch['loss_mask'].to(self.device)
1055
+ domain = batch.get('domain', 'bpe')
1056
+ task_type = batch.get('task_type', 'unknown')
1057
+
1058
+ # v26: NO autocast — full fp32 forward pass
1059
+ if True: # Flat block, no autocast context manager
1060
+ # 1. Embed
1061
+ x = self.multi_embed.embed(input_ids, domain)
1062
+
1063
+ # 2. Domain membrane
1064
+ x = self.plug_membranes(x, domain)
1065
+
1066
+ # 3. Memory read
1067
+ x = self.model.memory.read(x)
1068
+
1069
+ # 4. Transformer layers
1070
+ B, S = input_ids.shape
1071
+ mask = self.model.causal_mask[:, :, :S, :S].to(self.device)
1072
+ layer_outputs = []
1073
+ for layer in self.model.layers:
1074
+ x = layer(x, mask=mask)
1075
+ layer_outputs.append(x)
1076
+
1077
+ # 5. Final norm
1078
+ x = self.model.final_norm(x)
1079
+
1080
+ # 6. Router + specialists
1081
+ pooled = x.mean(dim=1)
1082
+ balance_loss = self.model.router.balance_loss(pooled)
1083
+
1084
+ # 7. Project to logits (domain-specific head)
1085
+ logits = self.multi_embed.project(x, domain)
1086
+
1087
+ # === Loss computation (C1 FIX: NO SHIFT — data already shifted) ===
1088
+ min_len = min(logits.shape[1], target_ids.shape[1], loss_mask.shape[1])
1089
+ flat_logits = logits[:, :min_len, :].reshape(-1, logits.size(-1))
1090
+ flat_targets = target_ids[:, :min_len].reshape(-1)
1091
+ flat_mask = loss_mask[:, :min_len].reshape(-1)
1092
+
1093
+ loss_per_token = F.cross_entropy(flat_logits, flat_targets, reduction='none', ignore_index=0)
1094
+ masked_loss = (loss_per_token * flat_mask).sum()
1095
+ num_tokens = flat_mask.sum().clamp(min=1)
1096
+ task_loss = masked_loss / num_tokens
1097
+
1098
+ # v26: no clamping — fp32 has full dynamic range
1099
+
1100
+ # L7 Auxiliary loss
1101
+ aux_loss = torch.tensor(0.0, device=self.device)
1102
+ l7_hidden = self._layer7_output
1103
+ self._layer7_output = None
1104
+ if l7_hidden is not None and domain == 'bpe':
1105
+ aux_logits = self.aux_head(l7_hidden)
1106
+ min_aux = min(aux_logits.shape[1], target_ids.shape[1])
1107
+ aux_flat = F.cross_entropy(
1108
+ aux_logits[:, :min_aux, :].reshape(-1, aux_logits.size(-1)),
1109
+ target_ids[:, :min_aux].reshape(-1),
1110
+ reduction='none', ignore_index=0,
1111
+ )
1112
+ aux_mask = loss_mask[:, :min_aux].reshape(-1)
1113
+ aux_loss = (aux_flat * aux_mask).sum() / aux_mask.sum().clamp(min=1)
1114
+
1115
+ # Gaussian loss (procedural — every N steps)
1116
+ gaussian_loss = torch.tensor(0.0, device=self.device)
1117
+ if task_type == 'gaussian' and len(layer_outputs) > 0:
1118
+ try:
1119
+ gs_out = self.gaussian_head(layer_outputs)
1120
+ # Simple anchor position loss against random targets
1121
+ target_pos = torch.randn_like(gs_out['anchors'][:, :, :3]) * self.gaussian_config.scene_scale
1122
+ gaussian_loss = F.mse_loss(gs_out['anchors'][:, :, :3], target_pos) * 0.1
1123
+ except Exception:
1124
+ pass # Non-critical
1125
+
1126
+ total_loss = (
1127
+ task_loss
1128
+ + self.cfg.balance_loss_weight * balance_loss
1129
+ + self.cfg.aux_loss_weight * aux_loss
1130
+ + gaussian_loss
1131
+ )
1132
+
1133
+ loss_dict = {
1134
+ 'task_loss': task_loss.item(),
1135
+ 'balance': balance_loss.item() if isinstance(balance_loss, torch.Tensor) else 0.0,
1136
+ 'aux_l7': aux_loss.item() if isinstance(aux_loss, torch.Tensor) else 0.0,
1137
+ 'gaussian': gaussian_loss.item() if isinstance(gaussian_loss, torch.Tensor) else 0.0,
1138
+ 'total': total_loss.item(),
1139
+ 'num_tokens': num_tokens.item(),
1140
+ 'task_type': task_type,
1141
+ 'domain': domain,
1142
+ }
1143
+ return total_loss, loss_dict
1144
+
1145
+ def save_checkpoint(self, path: str, session_mgr: KaggleSessionManager = None):
1146
+ """Save full training state."""
1147
+ # Check disk space
1148
+ free = disk_free_gb(str(Path(path).parent))
1149
+ if free < self.cfg.min_disk_gb:
1150
+ print(f" WARNING: Only {free:.1f} GB free, skipping checkpoint save")
1151
+ return False
1152
+
1153
+ state = {
1154
+ 'model_state_dict': self.model.state_dict(),
1155
+ 'multi_embed_state_dict': self.multi_embed.state_dict(),
1156
+ 'plug_membranes_state_dict': self.plug_membranes.state_dict(),
1157
+ 'aux_head_state_dict': self.aux_head.state_dict(),
1158
+ 'gaussian_head_state_dict': self.gaussian_head.state_dict(),
1159
+ 'optimizer_state_dict': self.optimizer.state_dict(),
1160
+ 'scaler_state_dict': self.scaler.state_dict() if self.scaler is not None else None,
1161
+ 'config': dataclasses.asdict(self.cfg),
1162
+ 'step': self.step,
1163
+ 'best_loss': self.best_loss,
1164
+ 'version': 'v27-kaggle-final',
1165
+ 'timestamp': datetime.datetime.now().isoformat(),
1166
+ 'rng_python': random.getstate(),
1167
+ 'rng_torch': torch.random.get_rng_state(),
1168
+ 'rng_cuda': torch.cuda.get_rng_state() if torch.cuda.is_available() else None,
1169
+ }
1170
+
1171
+ tmp = path + '.tmp'
1172
+ try:
1173
+ torch.save(state, tmp)
1174
+ # Verify
1175
+ verify = torch.load(tmp, map_location='cpu', weights_only=False)
1176
+ assert verify['step'] == self.step
1177
+ del verify
1178
+ if os.path.exists(path):
1179
+ os.remove(path)
1180
+ os.rename(tmp, path)
1181
+ print(f" Checkpoint saved: {path} (step {self.step})")
1182
+
1183
+ # Save manifest for session chaining
1184
+ if session_mgr is not None:
1185
+ session_mgr.save_manifest(self.step, os.path.basename(path))
1186
+
1187
+ return True
1188
+ except Exception as e:
1189
+ print(f" SAVE FAILED: {e}")
1190
+ if os.path.exists(tmp):
1191
+ os.remove(tmp)
1192
+ return False
1193
+
1194
+ def resume(self, path: str) -> bool:
1195
+ """Resume from checkpoint."""
1196
+ print(f"Resuming from: {path}")
1197
+ try:
1198
+ cp = torch.load(path, map_location='cpu', weights_only=False)
1199
+ except Exception as e:
1200
+ print(f" Failed to load checkpoint: {e}")
1201
+ return False
1202
+
1203
+ if 'model_state_dict' not in cp:
1204
+ print(" Invalid checkpoint (no model_state_dict)")
1205
+ return False
1206
+
1207
+ result = self.model.load_state_dict(cp['model_state_dict'], strict=False)
1208
+ if result.missing_keys:
1209
+ print(f" Missing keys: {len(result.missing_keys)}")
1210
+ if result.unexpected_keys:
1211
+ print(f" Unexpected keys: {len(result.unexpected_keys)}")
1212
+
1213
+ if 'multi_embed_state_dict' in cp:
1214
+ self.multi_embed.load_state_dict(cp['multi_embed_state_dict'], strict=False)
1215
+ if 'plug_membranes_state_dict' in cp:
1216
+ self.plug_membranes.load_state_dict(cp['plug_membranes_state_dict'], strict=False)
1217
+ if 'aux_head_state_dict' in cp:
1218
+ self.aux_head.load_state_dict(cp['aux_head_state_dict'], strict=False)
1219
+ if 'gaussian_head_state_dict' in cp:
1220
+ self.gaussian_head.load_state_dict(cp['gaussian_head_state_dict'], strict=False)
1221
+
1222
+ self.model.to(self.device)
1223
+ self.multi_embed.to(self.device)
1224
+ self.plug_membranes.to(self.device)
1225
+ self.aux_head.to(self.device)
1226
+ self.gaussian_head.to(self.device)
1227
+
1228
+ # Optimizer must be set up AFTER model is on device
1229
+ self.setup_optimizer()
1230
+
1231
+ if 'optimizer_state_dict' in cp:
1232
+ try:
1233
+ self.optimizer.load_state_dict(cp['optimizer_state_dict'])
1234
+ except: print(" Optimizer state load failed — using fresh")
1235
+
1236
+ # v26: no scaler to restore
1237
+
1238
+ self.step = cp.get('step', 0)
1239
+ self.best_loss = cp.get('best_loss', float('inf'))
1240
+
1241
+ # Restore RNG
1242
+ if 'rng_torch' in cp: torch.random.set_rng_state(cp['rng_torch'])
1243
+ if 'rng_cuda' in cp and cp['rng_cuda'] is not None and torch.cuda.is_available():
1244
+ torch.cuda.set_rng_state(cp['rng_cuda'])
1245
+ if 'rng_python' in cp: random.setstate(cp['rng_python'])
1246
+
1247
+ print(f" Resumed at step {self.step}, best_loss={self.best_loss:.4f}")
1248
+ return True
1249
+
1250
+ def train(self, sampler: UnifiedSampler, session_mgr: KaggleSessionManager):
1251
+ """Main training loop."""
1252
+ print(f"\n{'='*60}")
1253
+ print(f"WYRM Training — Steps {self.step} -> {self.cfg.total_steps}")
1254
+ print(f" Batch: {self.cfg.batch_size} x {self.cfg.grad_accumulation} = {self.cfg.batch_size * self.cfg.grad_accumulation}")
1255
+ print(f" VRAM: {vram_gb():.2f} / {vram_total_gb():.1f} GB")
1256
+ print(f" Disk: {disk_free_gb():.1f} GB free")
1257
+ print(f"{'='*60}\n")
1258
+
1259
+ accum = 0
1260
+ _shutdown = False
1261
+
1262
+ def _signal_handler(sig, frame):
1263
+ nonlocal _shutdown
1264
+ _shutdown = True
1265
+ print(f"\n Signal {sig} — saving and stopping...")
1266
+
1267
+ signal.signal(signal.SIGTERM, _signal_handler)
1268
+ signal.signal(signal.SIGINT, _signal_handler)
1269
+
1270
+ t_start = time.time()
1271
+
1272
+ while self.step < self.cfg.total_steps and not _shutdown:
1273
+ # ── Session time check ──
1274
+ if session_mgr.should_save_and_exit():
1275
+ print(f"\n Session time limit approaching ({session_mgr.elapsed_hours():.1f}h)")
1276
+ self.save_checkpoint(str(CKPT_OUT / 'latest.pt'), session_mgr)
1277
+ break
1278
+
1279
+ # ── Sample batch ──
1280
+ batch = sampler.sample_batch(self.cfg.batch_size, self.step)
1281
+
1282
+ try:
1283
+ total_loss, loss_dict = self.compute_loss(batch)
1284
+
1285
+ # ── NaN guard ──
1286
+ if not torch.isfinite(total_loss):
1287
+ self.nan_count += 1
1288
+ self.total_nan_count += 1
1289
+ print(f" Step {self.step}: NaN/Inf detected (retry {self.nan_count}/{self.cfg.max_nan_retries})")
1290
+
1291
+ # CRITICAL: clean VRAM on NaN (v24 bug #4 fix)
1292
+ self.optimizer.zero_grad(set_to_none=True)
1293
+ del total_loss, loss_dict, batch
1294
+ vram_cleanup()
1295
+ accum = 0
1296
+
1297
+ if self.nan_count >= self.cfg.max_nan_retries:
1298
+ print(f" Max NaN retries reached — skipping step")
1299
+ self.nan_count = 0
1300
+ self.step += 1 # Don't get stuck
1301
+
1302
+ if self.total_nan_count >= self.cfg.max_total_nans:
1303
+ print(f" FATAL: {self.total_nan_count} total NaNs — saving and stopping")
1304
+ self.save_checkpoint(str(CKPT_OUT / 'nan_rescue.pt'), session_mgr)
1305
+ break
1306
+ continue
1307
+
1308
+ self.nan_count = 0 # Reset per-step counter on success
1309
+
1310
+ # ── Backward (v26: direct, no scaler) ──
1311
+ scaled = total_loss / self.cfg.grad_accumulation
1312
+ scaled.backward()
1313
+
1314
+ except RuntimeError as e:
1315
+ err = str(e).lower()
1316
+ if 'out of memory' in err:
1317
+ print(f" OOM at step {self.step} — cleaning up")
1318
+ self.optimizer.zero_grad(set_to_none=True)
1319
+ vram_cleanup()
1320
+ accum = 0
1321
+ continue
1322
+ else:
1323
+ print(f" RuntimeError: {e}")
1324
+ self.save_checkpoint(str(CKPT_OUT / 'error_rescue.pt'), session_mgr)
1325
+ raise
1326
+
1327
+ accum += 1
1328
+
1329
+ if accum >= self.cfg.grad_accumulation:
1330
+ # ── Optimizer step (v26: direct, no scaler) ──
1331
+ all_params = (list(self.model.parameters()) +
1332
+ list(self.multi_embed.parameters()) +
1333
+ list(self.plug_membranes.parameters()) +
1334
+ list(self.aux_head.parameters()) +
1335
+ list(self.gaussian_head.parameters()))
1336
+ grad_norm = torch.nn.utils.clip_grad_norm_(all_params, self.cfg.max_grad_norm)
1337
+
1338
+ self.optimizer.step()
1339
+ self.optimizer.zero_grad(set_to_none=True)
1340
+ self.scheduler.step(self.step)
1341
+ self.step += 1
1342
+ accum = 0
1343
+
1344
+ # ── Update best loss ──
1345
+ avg_loss = loss_dict.get('total', 0.0)
1346
+ if avg_loss < self.best_loss:
1347
+ self.best_loss = avg_loss
1348
+
1349
+ # ── Logging ──
1350
+ elapsed = time.time() - t_start
1351
+ steps_per_sec = self.step / max(elapsed, 1)
1352
+ eta_h = (self.cfg.total_steps - self.step) / max(steps_per_sec, 0.001) / 3600
1353
+ phase = self.cfg.get_phase(self.step)
1354
+
1355
+ if self.step % self.cfg.log_every == 0 or DUMMY_MODE:
1356
+ mem = vram_gb()
1357
+ print(
1358
+ f"Step {self.step:>6}/{self.cfg.total_steps} | "
1359
+ f"{phase:11s} | {loss_dict.get('task_type','?'):12s} | "
1360
+ f"{loss_dict.get('domain','?'):4s} | "
1361
+ f"loss={avg_loss:.4f} | task={loss_dict.get('task_loss',0):.4f} | "
1362
+ f"grad={grad_norm:.3f} | VRAM={mem:.2f}GB | "
1363
+ f"best={self.best_loss:.4f} | ETA={eta_h:.1f}h"
1364
+ )
1365
+
1366
+ # ── VRAM monitoring ──
1367
+ if self.step % 10 == 0:
1368
+ mem = vram_gb()
1369
+ if mem > self.cfg.vram_warning_gb:
1370
+ print(f" WARNING: VRAM {mem:.2f}GB > {self.cfg.vram_warning_gb}GB — cleaning")
1371
+ vram_cleanup()
1372
+
1373
+ # ── Telemetry ──
1374
+ entry = {
1375
+ 'step': self.step, 'loss': avg_loss,
1376
+ 'task_loss': loss_dict.get('task_loss', 0),
1377
+ 'phase': phase, 'task': loss_dict.get('task_type', ''),
1378
+ 'domain': loss_dict.get('domain', ''),
1379
+ 'grad_norm': float(grad_norm),
1380
+ 'vram_gb': vram_gb(),
1381
+ 'elapsed_s': elapsed,
1382
+ 'timestamp': datetime.datetime.now().isoformat(),
1383
+ }
1384
+ self.telemetry.append(entry)
1385
+
1386
+ # ── Checkpoint ──
1387
+ if self.step % self.cfg.checkpoint_every == 0:
1388
+ self.save_checkpoint(str(CKPT_OUT / 'latest.pt'), session_mgr)
1389
+ # Write telemetry
1390
+ telem_path = TELEM / 'telemetry.jsonl'
1391
+ with open(telem_path, 'a') as f:
1392
+ for e in self.telemetry:
1393
+ f.write(json.dumps(e) + '\n')
1394
+ self.telemetry = []
1395
+
1396
+ # ── Phase transition ──
1397
+ for boundary, name in [(self.cfg.phase1_end, 'REASONING'),
1398
+ (self.cfg.phase2_end, 'DEPTH'),
1399
+ (self.cfg.phase3_end, 'OMEGA')]:
1400
+ if self.step == boundary:
1401
+ print(f"\n === PHASE TRANSITION: {name} ===\n")
1402
+
1403
+ # ── Final save ──
1404
+ if self.total_nan_count < self.cfg.max_total_nans:
1405
+ self.save_checkpoint(str(CKPT_OUT / 'latest.pt'), session_mgr)
1406
+ self.save_checkpoint(str(RUNS / f'wyrm_step_{self.step}.pt'))
1407
+
1408
+ # ── Write final telemetry ──
1409
+ if self.telemetry:
1410
+ with open(TELEM / 'telemetry.jsonl', 'a') as f:
1411
+ for e in self.telemetry:
1412
+ f.write(json.dumps(e) + '\n')
1413
+
1414
+ # ── Generate HTML dashboard ──
1415
+ self._write_dashboard()
1416
+
1417
+ print(f"\nTraining complete at step {self.step}. Best loss: {self.best_loss:.4f}")
1418
+ print(f" Total NaNs: {self.total_nan_count}")
1419
+ vram_peak = torch.cuda.max_memory_allocated()/(1024**3) if torch.cuda.is_available() else 0.0
1420
+ print(f" VRAM peak: {vram_peak:.2f} GB")
1421
+
1422
+ def _write_dashboard(self):
1423
+ """Generate a simple HTML telemetry dashboard."""
1424
+ telem_path = TELEM / 'telemetry.jsonl'
1425
+ if not telem_path.exists():
1426
+ return
1427
+
1428
+ entries = []
1429
+ with open(telem_path) as f:
1430
+ for line in f:
1431
+ if line.strip():
1432
+ entries.append(json.loads(line))
1433
+
1434
+ if not entries:
1435
+ return
1436
+
1437
+ steps = [e['step'] for e in entries]
1438
+ losses = [e['loss'] for e in entries]
1439
+ vrams = [e.get('vram_gb', 0) for e in entries]
1440
+
1441
+ html = f"""<!DOCTYPE html>
1442
+ <html><head><title>WYRM v25 Training Dashboard</title>
1443
+ <style>
1444
+ body {{ font-family: monospace; background: #1a1a2e; color: #eee; padding: 20px; }}
1445
+ .card {{ background: #16213e; padding: 15px; margin: 10px 0; border-radius: 8px; }}
1446
+ h1 {{ color: #e94560; }}
1447
+ .stat {{ display: inline-block; margin: 0 20px; }}
1448
+ .stat-val {{ font-size: 24px; color: #0f3460; font-weight: bold; }}
1449
+ </style></head><body>
1450
+ <h1>WYRM 500M Training Dashboard (v25)</h1>
1451
+ <div class="card">
1452
+ <div class="stat"><div>Steps</div><div class="stat-val">{self.step}</div></div>
1453
+ <div class="stat"><div>Best Loss</div><div class="stat-val">{self.best_loss:.4f}</div></div>
1454
+ <div class="stat"><div>Total NaN</div><div class="stat-val">{self.total_nan_count}</div></div>
1455
+ <div class="stat"><div>VRAM Peak</div><div class="stat-val">{torch.cuda.max_memory_allocated()/(1024**3) if torch.cuda.is_available() else 0.0:.2f} GB</div></div>
1456
+ </div>
1457
+ <div class="card">
1458
+ <h2>Loss Curve</h2>
1459
+ <pre>"""
1460
+
1461
+ # ASCII loss chart
1462
+ if losses:
1463
+ max_loss = max(losses[:10]) if len(losses) >= 10 else max(losses) # Scale to early loss
1464
+ min_loss = min(losses)
1465
+ chart_width = 80
1466
+ chart_height = 20
1467
+ for row in range(chart_height):
1468
+ threshold = max_loss - (max_loss - min_loss) * row / chart_height
1469
+ line = f"{threshold:7.3f} |"
1470
+ step_stride = max(1, len(losses) // chart_width)
1471
+ for i in range(0, min(len(losses), chart_width * step_stride), step_stride):
1472
+ line += "#" if losses[i] >= threshold else " "
1473
+ html += line + "\n"
1474
+ html += f" +{''.join(['-'] * min(chart_width, len(losses)))}\n"
1475
+ html += f" Step 0{' ' * (min(chart_width, len(losses)) - 10)}Step {steps[-1]}\n"
1476
+
1477
+ html += f"""</pre>
1478
+ </div>
1479
+ <div class="card">
1480
+ <h2>Training Log (last 20 steps)</h2>
1481
+ <pre>"""
1482
+ for e in entries[-20:]:
1483
+ html += (f"Step {e['step']:>6} | loss={e['loss']:.4f} | "
1484
+ f"task={e.get('task_loss',0):.4f} | "
1485
+ f"VRAM={e.get('vram_gb',0):.2f}GB | "
1486
+ f"{e.get('phase','?')} | {e.get('task','?')}\n")
1487
+
1488
+ html += f"""</pre>
1489
+ </div>
1490
+ <div class="card">
1491
+ <p>Generated: {datetime.datetime.now().isoformat()}</p>
1492
+ <p>GPU: {GPU_NAME} | Dtype: float32 (FULL) | GradScaler: OFF</p>
1493
+ </div>
1494
+ </body></html>"""
1495
+
1496
+ dashboard_path = OUTPUT / 'dashboard.html'
1497
+ dashboard_path.write_text(html)
1498
+ print(f" Dashboard: {dashboard_path}")
1499
+
1500
+
1501
+ # ═══════════════════════════════════════════════════════════════
1502
+ # 12. MAIN
1503
+ # ═══════════════════════════════════════════════════════════════
1504
+
1505
+ def main():
1506
+ # ── Config ──
1507
+ if DUMMY_MODE:
1508
+ cfg = make_dummy_config()
1509
+ else:
1510
+ cfg = WyrmKaggleConfig()
1511
+
1512
+ # ── Seed ──
1513
+ torch.manual_seed(cfg.seed)
1514
+ random.seed(cfg.seed)
1515
+ if torch.cuda.is_available():
1516
+ torch.cuda.manual_seed(cfg.seed)
1517
+ torch.backends.cudnn.benchmark = True
1518
+
1519
+ device = DEVICE # Use global DEVICE (cuda or cpu for dummy)
1520
+
1521
+ # ── Session manager ──
1522
+ session_mgr = KaggleSessionManager(OUTPUT, CKPT_INPUT)
1523
+
1524
+ # ── Build trainer ──
1525
+ trainer = WyrmKaggleTrainer(cfg, device)
1526
+
1527
+ # ── Resume if checkpoint found ──
1528
+ if session_mgr.resume_path:
1529
+ if not trainer.resume(session_mgr.resume_path):
1530
+ print(" Resume failed — starting fresh")
1531
+ trainer.setup_optimizer()
1532
+ else:
1533
+ trainer.setup_optimizer()
1534
+
1535
+ # ── Load tokenizer ──
1536
+ bpe_tokenizer = None
1537
+ tokenizer_paths = [
1538
+ WORKING / 'gladius_bpe_16k.model', # Copied during setup
1539
+ INPUT_BASE / 'gladius_bpe_16k.model',
1540
+ INPUT_BASE / 'data' / 'gladius_bpe_16k.model',
1541
+ INPUT_BASE / 'data' / 'gladius_multimodal_32k.model',
1542
+ INPUT_BASE / 'gladius_multimodal_32k.model',
1543
+ Path('/kaggle/input/wyrm-500m-base/gladius_bpe_16k.model'),
1544
+ Path('/kaggle/input/wyrm-500m-base/data/gladius_multimodal_32k.model'),
1545
+ Path('/kaggle/input/datasets/avashakil/wyrm-500m-base/gladius_bpe_16k.model'),
1546
+ ]
1547
+ for tp in tokenizer_paths:
1548
+ if tp.exists():
1549
+ bpe_tokenizer = load_bpe_tokenizer(str(tp))
1550
+ break
1551
+
1552
+ if bpe_tokenizer is None and not DUMMY_MODE:
1553
+ print(" WARNING: No BPE tokenizer found — using procedural data only")
1554
+
1555
+ # ── Build data loaders ──
1556
+ text_ds = ProceduralTextDataset(bpe_tokenizer, cfg.max_seq_len)
1557
+ math_ds = ProceduralMathDataset(cfg.max_seq_len)
1558
+ gauss_ds = ProceduralGaussianDataset(
1559
+ size=1000, gaussians_per_object=32 if DUMMY_MODE else 128,
1560
+ max_objects=3 if DUMMY_MODE else 5, device='cpu')
1561
+
1562
+ # TODO: Load real corpus files from dataset when available
1563
+ # For now, procedural data proves the training pipeline works end-to-end
1564
+ if not DUMMY_MODE and IS_KAGGLE:
1565
+ # Try to load real corpus data
1566
+ data_dir = INPUT_BASE / 'data'
1567
+ if (data_dir / 'cognition').exists():
1568
+ print(f" Real data available at {data_dir}")
1569
+ # Future: wire CognitiveCorpus, TextCorpus etc. here
1570
+
1571
+ sampler = UnifiedSampler(cfg, text_ds, math_ds, gauss_ds)
1572
+
1573
+ # ── Train ──
1574
+ trainer.train(sampler, session_mgr)
1575
+
1576
+ print("\n" + "=" * 60)
1577
+ print("WYRM v25 — Training Complete")
1578
+ print("=" * 60)
1579
+
1580
+
1581
+ if __name__ == '__main__':
1582
+ main()