WYRM v27 FINAL notebook — training on Kaggle T4 x2
Browse files- 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()
|