File size: 115,482 Bytes
11f07f9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 1064 1065 1066 1067 1068 1069 1070 1071 1072 1073 1074 1075 1076 1077 1078 1079 1080 1081 1082 1083 1084 1085 1086 1087 1088 1089 1090 1091 1092 1093 1094 1095 1096 1097 1098 1099 1100 1101 1102 1103 1104 1105 1106 1107 1108 1109 1110 1111 1112 1113 1114 1115 1116 1117 1118 1119 1120 1121 1122 1123 1124 1125 1126 1127 1128 1129 1130 1131 1132 1133 1134 1135 1136 1137 1138 1139 1140 1141 1142 1143 1144 1145 1146 1147 1148 1149 1150 1151 1152 1153 1154 1155 1156 1157 1158 1159 1160 1161 1162 1163 1164 1165 1166 1167 1168 1169 1170 1171 1172 1173 1174 1175 1176 1177 1178 1179 1180 1181 1182 1183 1184 1185 1186 1187 1188 1189 1190 1191 1192 1193 1194 1195 1196 1197 1198 1199 1200 1201 1202 1203 1204 1205 1206 1207 1208 1209 1210 1211 1212 1213 1214 1215 1216 1217 1218 1219 1220 1221 1222 1223 1224 1225 1226 1227 1228 1229 1230 1231 1232 1233 1234 1235 1236 1237 1238 1239 1240 1241 1242 1243 1244 1245 1246 1247 1248 1249 1250 1251 1252 1253 1254 1255 1256 1257 1258 1259 1260 1261 1262 1263 1264 1265 1266 1267 1268 1269 1270 1271 1272 1273 1274 1275 1276 1277 1278 1279 1280 1281 1282 1283 1284 1285 1286 1287 1288 1289 1290 1291 1292 1293 1294 1295 1296 1297 1298 1299 1300 1301 1302 1303 1304 1305 1306 1307 1308 1309 1310 1311 1312 1313 1314 1315 1316 1317 1318 1319 1320 1321 1322 1323 1324 1325 1326 1327 1328 1329 1330 1331 1332 1333 1334 1335 1336 1337 1338 1339 1340 1341 1342 1343 1344 1345 1346 1347 1348 1349 1350 1351 1352 1353 1354 1355 1356 1357 1358 1359 1360 1361 1362 1363 1364 1365 1366 1367 1368 1369 1370 1371 1372 1373 1374 1375 1376 1377 1378 1379 1380 1381 1382 1383 1384 1385 1386 1387 1388 1389 1390 1391 1392 1393 1394 1395 1396 1397 1398 1399 1400 1401 1402 1403 1404 1405 1406 1407 1408 1409 1410 1411 1412 1413 1414 1415 1416 1417 1418 1419 1420 1421 1422 1423 1424 1425 1426 1427 1428 1429 1430 1431 1432 1433 1434 1435 1436 1437 1438 1439 1440 1441 1442 1443 1444 1445 1446 1447 1448 1449 1450 1451 1452 1453 1454 1455 1456 1457 1458 1459 1460 1461 1462 1463 1464 1465 1466 1467 1468 1469 1470 1471 1472 1473 1474 1475 1476 1477 1478 1479 1480 1481 1482 1483 1484 1485 1486 1487 1488 1489 1490 1491 1492 1493 1494 1495 1496 1497 1498 1499 1500 1501 1502 1503 1504 1505 1506 1507 1508 1509 1510 1511 1512 1513 1514 1515 1516 1517 1518 1519 1520 1521 1522 1523 1524 1525 1526 1527 1528 1529 1530 1531 1532 1533 1534 1535 1536 1537 1538 1539 1540 1541 1542 1543 1544 1545 1546 1547 1548 1549 1550 1551 1552 1553 1554 1555 1556 1557 1558 1559 1560 1561 1562 1563 1564 1565 1566 1567 1568 1569 1570 1571 1572 1573 1574 1575 1576 1577 1578 1579 1580 1581 1582 1583 1584 1585 1586 1587 1588 1589 1590 1591 1592 1593 1594 1595 1596 1597 1598 1599 1600 1601 1602 1603 1604 1605 1606 1607 1608 1609 1610 1611 1612 1613 1614 1615 1616 1617 1618 1619 1620 1621 1622 1623 1624 1625 1626 1627 1628 1629 1630 1631 1632 1633 1634 1635 1636 1637 1638 1639 1640 1641 1642 1643 1644 1645 1646 1647 1648 1649 1650 1651 1652 1653 1654 1655 1656 1657 1658 1659 1660 1661 1662 1663 1664 1665 1666 1667 1668 1669 1670 1671 1672 1673 1674 1675 1676 1677 1678 1679 1680 1681 1682 1683 1684 1685 1686 1687 1688 1689 1690 1691 1692 1693 1694 1695 1696 1697 1698 1699 1700 1701 1702 1703 1704 1705 1706 1707 1708 1709 1710 1711 1712 1713 1714 1715 1716 1717 1718 1719 1720 1721 1722 1723 1724 1725 1726 1727 1728 1729 1730 1731 1732 1733 1734 1735 1736 1737 1738 1739 1740 1741 1742 1743 1744 1745 1746 1747 1748 1749 1750 1751 1752 1753 1754 1755 1756 1757 1758 1759 1760 1761 1762 1763 1764 1765 1766 1767 1768 1769 1770 1771 1772 1773 1774 1775 1776 1777 1778 1779 1780 1781 1782 1783 1784 1785 1786 1787 1788 1789 1790 1791 1792 1793 1794 1795 1796 1797 1798 1799 1800 1801 1802 1803 1804 1805 1806 1807 1808 1809 1810 1811 1812 1813 1814 1815 1816 1817 1818 1819 1820 1821 1822 1823 1824 1825 1826 1827 1828 1829 1830 1831 1832 1833 1834 1835 1836 1837 1838 1839 1840 1841 1842 1843 1844 1845 1846 1847 1848 1849 1850 1851 1852 1853 1854 1855 1856 1857 1858 1859 1860 1861 1862 1863 1864 1865 1866 1867 1868 1869 1870 1871 1872 1873 1874 1875 1876 1877 1878 1879 1880 1881 1882 1883 1884 1885 1886 1887 1888 1889 1890 1891 1892 1893 1894 1895 1896 1897 1898 1899 1900 1901 1902 1903 1904 1905 1906 1907 1908 1909 1910 1911 1912 1913 1914 1915 1916 1917 1918 1919 1920 1921 1922 1923 1924 1925 1926 1927 1928 1929 1930 1931 1932 1933 1934 1935 1936 1937 1938 1939 1940 1941 1942 1943 1944 1945 1946 1947 1948 1949 1950 1951 1952 1953 1954 1955 1956 1957 1958 1959 1960 1961 1962 1963 1964 1965 1966 1967 1968 1969 1970 1971 1972 1973 1974 1975 1976 1977 1978 1979 1980 1981 1982 1983 1984 1985 1986 1987 1988 1989 1990 1991 1992 1993 1994 1995 1996 1997 1998 1999 2000 2001 2002 2003 2004 2005 2006 2007 2008 2009 2010 2011 2012 2013 2014 2015 2016 2017 2018 2019 2020 2021 2022 2023 2024 2025 2026 2027 2028 2029 2030 2031 2032 2033 2034 2035 2036 2037 2038 2039 2040 2041 2042 2043 2044 2045 2046 2047 2048 2049 2050 2051 2052 2053 2054 2055 2056 2057 2058 2059 2060 2061 2062 2063 2064 2065 2066 2067 2068 2069 2070 2071 2072 2073 2074 2075 2076 2077 2078 2079 2080 2081 2082 2083 2084 2085 2086 2087 2088 2089 2090 2091 2092 2093 2094 2095 2096 2097 2098 2099 2100 2101 2102 2103 2104 2105 2106 2107 2108 2109 2110 2111 2112 2113 2114 2115 2116 2117 2118 2119 2120 2121 2122 2123 2124 2125 2126 2127 2128 2129 2130 2131 2132 2133 2134 2135 2136 2137 2138 2139 2140 2141 2142 2143 2144 2145 2146 2147 2148 2149 2150 2151 2152 2153 2154 2155 2156 2157 2158 2159 2160 2161 2162 2163 2164 2165 2166 2167 2168 2169 2170 2171 2172 2173 2174 2175 2176 2177 2178 2179 2180 2181 2182 2183 2184 2185 2186 2187 2188 2189 2190 2191 2192 2193 2194 2195 2196 2197 2198 2199 2200 2201 2202 2203 2204 2205 2206 2207 2208 2209 2210 2211 2212 2213 2214 2215 2216 2217 2218 2219 2220 2221 2222 2223 2224 2225 2226 2227 2228 2229 2230 2231 2232 2233 2234 2235 2236 2237 2238 2239 2240 2241 2242 2243 2244 2245 2246 2247 2248 2249 2250 2251 2252 2253 2254 2255 2256 2257 2258 2259 2260 2261 2262 2263 2264 2265 2266 2267 2268 2269 2270 2271 2272 2273 2274 2275 2276 2277 2278 2279 2280 2281 2282 2283 2284 2285 2286 2287 2288 2289 2290 2291 2292 2293 2294 2295 2296 2297 2298 2299 2300 2301 2302 2303 2304 2305 2306 2307 2308 2309 2310 2311 2312 2313 2314 2315 2316 2317 2318 2319 2320 2321 2322 2323 2324 2325 2326 2327 2328 2329 2330 2331 2332 2333 2334 2335 2336 2337 2338 2339 2340 2341 2342 2343 2344 2345 2346 2347 2348 2349 2350 2351 2352 2353 2354 2355 2356 2357 2358 2359 2360 2361 2362 2363 2364 2365 2366 2367 2368 2369 2370 | # decoderstack_medium_pt-sft-fable.py
#
# Single-file d24 pre-training pipeline with a handwritten forward/backward and a
# written-out optimizer: no autograd, no torch.optim, no param groups, no nn.Module.
#
# (From Chris -- Core design decisions):
# - No nn.Module, no m.to, no state_dict / load_state_dict.
# - Every tensor is created directly on the device, at its final dtype.
# - No accommodations for "prior checkpoints", we're starting from scratch.
# - No torch.optim or autograd, we're doing everything manually.
# - Use globals -- global cfg, global m -- don't pass things around.
# - The model is a plain class used as a namespace of plain torch.Tensors.
# nn.Parameter does nothing for us: Parameter exists for autograd leaf
# bookkeeping and Module registration, neither of which we use. Plain
# tensors are directly usable in the math (m.W_in, not m.W_in.weight),
# accept attached state (.grad32, .mantissa, ...) just like Parameters,
# and default to requires_grad=False -- which is what we want everywhere,
# because we implement grad.
# - Dtypes are hardcoded everywhere -- stated at creation, never inferred by
# matching another tensor's dtype. (No fp64 parity tier in this file.)
# - Hardcoded to the d24 config; none of nanochat's auto-scaling by model size.
# - Multi-GPU shards the optimizer, not the model (nanochat's scheme): every
# rank holds the full bf16 live weights and full grad accumulators, optimizer
# state is allocated at shard sizes, and optimizer_step wraps the same update
# kernels in reduce-scatter -> owned-shard update -> live all-gather.
# - We're not doing FP8 yet.
#
# The "§" technique defines the code sections in here.
#
# The model/training code comes from the nanochat repo, branch fwd-bwd
# (nanochat/train_step.py, nanochat/gpt.py). That branch's d24 run is the
# reference implementation we want to match -- we're refactoring and dropping
# baggage, not changing the math:
# C:\Users\chris\Documents\GitHub\agent-ops\nanochat\2026-07-29_0833am_d24-throughput-gap\NOTES.md
#
# The code below the seam (marked near the bottom) comes from the 'stacks' repo,
# pulled mainly for the pre-tokenized data + distributed loader and CORE eval.
#
# One-off derived quantities (parameter counts, flops/token, the training
# horizon, the LR/WD batch corrections, cu_seqlens sizing) are HARDCODED in
# this script; `scaling.py` (kept alongside it) recomputes and documents them.
# --------------------------------------------------------------------------------
# § Setup
# --------------------------------------------------------------------------------
import os
import sys
import time as _time
run_wall_t0 = _time.perf_counter()
del _time
with open(sys.argv[0], 'r') as f:
code = f.read() # the run section logs the script source to wandb
import datetime
import gc
import glob
import json
import math
import random
import threading
import time
from pathlib import Path
from types import SimpleNamespace
from typing import NamedTuple
import numpy as np
import wandb
os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True"
os.environ["HF_HUB_DISABLE_PROGRESS_BARS"] = "1"
import torch
import torch._dynamo as dynamo
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from kernels import get_kernel
dynamo.config.recompile_limit = 64
# ==== Distributed setup ====
# dist is always initialized (launch under torchrun, even for one process) --
# the data pipeline below the seam uses dist.barrier() and the loader shards
# by rank.
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
assert torch.cuda.is_available()
device = torch.device("cuda", int(os.environ["LOCAL_RANK"]))
torch.cuda.set_device(device)
dist.init_process_group(backend="nccl", device_id=device)
dist.barrier()
master_process = (rank == 0)
def print0(*args, console=False, **kwargs):
if master_process:
print(*args, **kwargs)
# ==== Wandb helpers ====
class DummyWandb:
"""No-op wandb replacement when logging is disabled."""
def log(self, *args, **kwargs): pass
def finish(self): pass
# BF16 dense peak FLOPS by GPU, for the MFU denominator. Just the GPUs this
# pipeline actually runs on; the full many-vendor table (and sources) lives in
# scaling.py. GH200 carries the same H100-class SXM die: 989 TFLOPS.
PEAK_FLOPS = {"GH200": 989e12, "H100": 989e12, "A100": 312e12}
def next_multiple_of_n(v: float | int, *, n: int):
return next(x for x in range(n, int(v) + 1 + n, n) if x >= v)
# --------------------------------------------------------------------------------
# § Flash Attention (raw FA3 forward/backward)
# --------------------------------------------------------------------------------
# The handwritten backward calls FA3's raw _flash_attn_forward/_flash_attn_backward
# torch.library ops directly -- no autograd Function in between. The forward
# returns the softmax LSE, which the backward consumes alongside the stashed
# output. FA3 only; there is no SDPA/naive fallback in this file.
_cc_major, _ = torch.cuda.get_device_capability()
if _cc_major == 9: # Hopper: the varunneal build gets better H100 results
fa3 = get_kernel("varunneal/flash-attention-3").flash_attn_interface
RAW_BWD_TAKES_BUFFERS = False # raw backward allocates and RETURNS dq/dk/dv
else: # Ampere sm80/86 / Ada sm89: community FA3 build
assert _cc_major == 8, f"FA3 required (sm8x or sm90); got sm{_cc_major}x"
_k = get_kernel("kernels-community/flash-attn3")
# The raw ops live in flash_attn_interface; the top level only re-exports
# the varlen/kvcache wrappers.
fa3 = getattr(_k, "flash_attn_interface", _k)
RAW_BWD_TAKES_BUFFERS = True # raw backward takes pre-allocated dq/dk/dv buffers
def flash_attn_varlen_fwd_lse(q, k, v, cu_seqlens, max_seqlen, window_size):
"""Attention forward that also returns what the handwritten backward needs:
(out, softmax_lse), with lse (H, T) fp32."""
out, softmax_lse, *_ = fa3._flash_attn_forward(
q, k, v,
cu_seqlens_q=cu_seqlens, cu_seqlens_k=cu_seqlens,
max_seqlen_q=max_seqlen, max_seqlen_k=max_seqlen,
softmax_scale=q.shape[-1] ** -0.5, causal=True,
window_size_left=window_size[0], window_size_right=window_size[1])
return out, softmax_lse
def flash_attn_varlen_bwd(dout, q, k, v, out, softmax_lse, cu_seqlens, max_seqlen, window_size):
"""Attention backward for flash_attn_varlen_fwd_lse: returns (dq, dk, dv).
The two FA3 builds' raw backward ops differ in calling convention -- the
sm80 community build's schema takes pre-allocated dq/dk/dv buffers (grads
come back through them), the sm90 varunneal build's allocates and returns
them -- hence the branch on the module-level flag."""
softmax_scale = q.shape[-1] ** -0.5
if RAW_BWD_TAKES_BUFFERS:
dq, dk, dv = torch.empty_like(q), torch.empty_like(k), torch.empty_like(v)
fa3._flash_attn_backward(
dout, q, k, v, out, softmax_lse,
cu_seqlens, cu_seqlens, # cu_seqlens_q, cu_seqlens_k
None, None, # seqused_q, seqused_k
max_seqlen, max_seqlen,
dq, dk, dv,
softmax_scale,
True, # is_causal
window_size[0], window_size[1],
0.0, # softcap
False, # deterministic
0, # sm_margin
)
else:
dq, dk, dv, _ = fa3._flash_attn_backward(
dout, q, k, v, out, softmax_lse,
cu_seqlens, cu_seqlens, # cu_seqlens_q, cu_seqlens_k
None, None, # seqused_q, seqused_k
max_seqlen, max_seqlen,
softmax_scale,
True, # is_causal
window_size[0], window_size[1],
0.0, # softcap
False, # deterministic
0, # sm_margin
)
return dq, dk, dv
# --------------------------------------------------------------------------------
# § Model Config
# --------------------------------------------------------------------------------
# Value embeddings (ResFormer-style) live on alternating layers, last always
# included. Banked over just the VE layers; ve_index maps layer -> bank slot
# (-1 = no VE on this layer) and is read by every forward body.
#
# Note: Deriving head size or count from d_model is a bad habit that has
# propagated through ~everyone's model code.
# There are only three real constraints--these values must match:
# 1. Number of key and value heads
# 2. Query-key head sizes
# 3. Value-output head sizes
#
# Recommended short window size:
# -(-seq_len // 4 // 128) * 128 # ceil to FA3 tile size (2048 -> 768)
class StackConfig:
# Model
n_layers: int = 24
d_model: int = 1536
# Input
d_vocab: int = 32768 # Must arrive padded (tensor cores, sharding) -- no
# auto-padding in this file; asserted below.
d_smr_gate: int = 24 # Input to smear gate is first 'd' positions of the
# normed input embedding.
# Attention
n_q_heads: int = 12
n_kv_heads: int = 12
n_o_heads: int = 12 # TODO - fold into n_qo_heads, since the code doesn't support
# a different ratio (group size) for qk vs. vo.
d_qk: int = 128 # Note: FA2 requires d_qk == d_vo, FA3 does not.
d_vo: int = 128
# Context and Sliding Window Attention
seq_len: int = 2048
short_win_size: int = 768
full_ctxt_layers: list[int] = [ 3, 7, 11, 15, 19, 23] # "sssL" pattern
window_sizes: list[tuple[int, int]] # Derived below.
# Attention - Value Embeddings
d_ve_gate: int = 12 # First 'd' positions of residual stream (after x0
# blending and norm) are the gate input.
# ve gates exist per head, per layer.
ve_layers: list[int] = [1, 3, 5, 7, 9, 11, 13, 15, 17, 19, 21, 23] # 0-indexed.
ve_index: list[int] # Derived from ve_layers.
num_ves: int
# MLP
d_mlp: int = 4 * 1536
# Training batch (nanochat d24 speedrun spec). Tokens, not sequences: with
# varlen packing a micro-batch is one packed 1-D stream, so the token count
# is the real quantity (= 16 seqs x 2048 in nanochat's batched terms).
# Total batch 2^20 tokens/step is nanochat's Power Lines auto-compute for
# d24.
micro_batch_tokens: int = 65536 # per rank, per micro-batch
total_batch_size: int = 2**20 # tokens per optimizer step
# Training horizon: the d24 speedrun spec (data:param ratio 8) --
# 8 x 729,810,624 scaling params = 5,838,484,992 tokens // 2^20 per step
# = 5,568 steps. Derivation: scaling.py.
num_iterations: int = 5568
# Evaluation and logging
val_tokens: int = 10485760 # per val-bpb pass: 320 training-shaped micro-batches
val_loss_every: int = 250
eval_buffer_tokens: int = 65536 # CORE/chat eval packing buffer. Eval is
# forward-only (no stash, no grads), so a
# buffer well past the training micro-batch
# fits easily; the rotary cache is sized to
# cover it.
save_checkpoint: bool = True
# Mid-run checkpoint capture, in COMPLETED optimizer steps: state is
# written on entering these loop steps (the final state always saves).
# 1950 = the first LR/momentum-cooldown step at the 5568-step horizon
# (the hold ends after update 1949 = N - round(0.65*N)): the last
# uncooled state -- the one to resume from to train the horizon longer.
save_steps: tuple = (1950,)
run_id: str = f"{str(datetime.datetime.now().strftime('%Y-%m-%d_%H%M%S'))}-d24"
wandb_run: str = "dummy" # "dummy" disables wandb
wandb_project: str = "decoderstack"
cfg = StackConfig() # Make config a global, don't pass it around.
# Sanity: the constraints the axes above must satisfy.
assert cfg.d_vocab % 64 == 0, "vocab must arrive padded to 64 (no auto-padding here)"
assert cfg.n_o_heads == cfg.n_q_heads, "attention output consumes one slot per query head"
assert cfg.n_q_heads % cfg.n_kv_heads == 0, "GQA needs query heads to tile over kv heads"
assert cfg.d_qk % 2 == 0, "rotary splits the qk head dim in half"
assert cfg.full_ctxt_layers[-1] == cfg.n_layers - 1, "final layer recommended to have full context"
# Derived quantities:
# Map layers to VE bank slots.
cfg.ve_index = [cfg.ve_layers.index(i) if i in cfg.ve_layers else -1 for i in range(cfg.n_layers)]
cfg.num_ves = len(cfg.ve_layers)
# Per-layer window sizes for sliding window attention.
# List of (left, right) tuples for FA3's window_size parameter:
# - left: how many tokens before current position to attend to
# - right: how many tokens after current position to attend to (0 for causal)
# "Full context" is (seq_len, 0): documents are at most seq_len tokens and
# varlen attention is doc-isolated, so a seq_len window is unlimited in effect.
cfg.window_sizes = [(cfg.short_win_size, 0)] * cfg.n_layers # All short, ...
for i in cfg.full_ctxt_layers:
cfg.window_sizes[i] = (cfg.seq_len, 0) # ... then overwrite with full.
# Derived batch quantities. Fixed total => grad accum scales down as GPUs are
# added: 32 at world=1, 4 at world=8. grad_scale rides into forward_backward
# as loss_scale, replacing the loss division of an autograd loop; at world>1
# it composes with ReduceOp.AVG grad comm to give the global batch mean.
assert cfg.total_batch_size % (cfg.micro_batch_tokens * world_size) == 0, \
"total batch must divide evenly into per-rank micro-batches"
grad_accum_steps = cfg.total_batch_size // (cfg.micro_batch_tokens * world_size)
grad_scale = 1 / grad_accum_steps
# --------------------------------------------------------------------------------
# § Shard Assignment
# --------------------------------------------------------------------------------
# Each GPU is responsible for a "shard" of the optimizer work:
# - Muon banks shard over their layer axis (dim 0).
# - AdamW params shard over the row axis of their (rows, cols) view -- vocab
# rows for input_embeds/lm_head, flattened (ve_slot * vocab) rows for
# value_embeds. (ve_slot alone is too small to divide across a world, and the
# rows are interchangeable for AdamW's elementwise update.)
# - ve_gate is NOT sharded: it is tiny (~thousands of floats) and ragged
# against world sizes, so every rank runs the full-size update instead.
# - grad32 always stays FULL size on every rank -- it is the source buffer for
# the reduce-scatter, not a shard.
#
# No zero-padding support: every sharded axis must divide evenly (asserted
# below). d24's axes -- 24 layers, 32768 vocab rows, 393,216 ve rows -- all
# divide by the world sizes we'd run (1, 2, 4, 8).
#
# At world_size == 1 every shard IS the whole tensor: the slices below span
# their full axes and optimizer_step's collectives short-circuit. One code
# path, degenerate comm.
assert cfg.n_layers % world_size == 0, \
f"Muon layer-sharding needs n_layers % world == 0 ({cfg.n_layers} % {world_size})"
layer_shard_size = cfg.n_layers // world_size
layer_shard_start = rank * layer_shard_size
layer_shard_slice = slice(layer_shard_start, layer_shard_start + layer_shard_size)
assert cfg.d_vocab % world_size == 0, \
f"AdamW row-sharding needs vocab % world == 0 ({cfg.d_vocab} % {world_size})"
vocab_shard_size = cfg.d_vocab // world_size
vocab_shard_start = rank * vocab_shard_size
vocab_shard_slice = slice(vocab_shard_start, vocab_shard_start + vocab_shard_size)
ve_rows = cfg.num_ves * cfg.d_vocab
assert ve_rows % world_size == 0, \
f"AdamW row-sharding needs ve_slot*vocab % world == 0 ({ve_rows} % {world_size})"
ve_row_shard_size = ve_rows // world_size
ve_row_shard_start = rank * ve_row_shard_size
ve_row_shard_slice = slice(ve_row_shard_start, ve_row_shard_start + ve_row_shard_size)
# --------------------------------------------------------------------------------
# § Model Initialization
# --------------------------------------------------------------------------------
class Model:
"""Namespace of plain tensors -- the live weights. Each weight also carries
its training state as attached attributes, allocated alongside it below:
.grad32 full-size gradient accumulator (fp32; bf16 for the two
embedding tables), explicitly zeroed between steps
.grad32_slices per-layer views of grad32 for the 3-D banks (see below)
.mantissa lower 16 bits of the fp32 master (uint16, shard-size)
.frst_mntm Muon first moment (fp32, shard-size)
.scnd_mntm Muon factored second moment (fp32, shard-size)
.residual_dim the weight axis that faces the residual stream (-1 or -2);
NorMuon's per-neuron mean-square is taken along it
.exp_avg AdamW first moment (fp32, shard-size)
.exp_avg_sq AdamW second moment (fp32, shard-size)
"""
# Input
input_embeds: Tensor
smear_gate: Tensor
smear_lambda: Tensor
# Attention
W_Q: Tensor
W_K: Tensor
W_V: Tensor
W_O: Tensor
value_embeds: Tensor
ve_gate: Tensor
# MLP
W_in: Tensor
W_out: Tensor
# Cross-Layer
resid_lambdas: Tensor # Per-layer gain on the residual stream.
x0_lambdas: Tensor # Per-layer coefficient for reading the input embedding.
backout_lambda: Tensor # How much of layer 16's output to remove from the stream
# prior to the lm head.
# Output
lm_head: Tensor
# Buffers (rotary cache; not trained, not checkpointed)
cos: Tensor
sin: Tensor
# The trained weights, in declaration order -- this tuple defines "every
# trained weight". __iter__ walks them so call sites can just say
# `for p in m` (grad zeroing); the names key the checkpoint dicts.
weight_names = ("input_embeds", "smear_gate", "smear_lambda",
"W_Q", "W_K", "W_V", "W_O", "value_embeds", "ve_gate",
"W_in", "W_out", "resid_lambdas", "x0_lambdas",
"backout_lambda", "lm_head")
def __iter__(self):
return (getattr(self, n) for n in self.weight_names)
# ==== Tensor Creation Idioms ====
# Reduce the boilerplate for defining weights and buffers.
fp32_empty = lambda *shape: torch.empty(*shape, dtype=torch.float32, device=device)
bf16_empty = lambda *shape: torch.empty(*shape, dtype=torch.bfloat16, device=device)
fp32_zeros = lambda *shape: torch.zeros(*shape, dtype=torch.float32, device=device)
bf16_zeros = lambda *shape: torch.zeros(*shape, dtype=torch.bfloat16, device=device)
uint16_zeros = lambda *shape: torch.zeros(*shape, dtype=torch.uint16, device=device)
# We use fp32 for the "master" weights, which are what we store on disk, and for
# avoiding rounding off small optimizer updates.
# All forward and backward computation is done on bf16 matrices (the "live" weights).
# Note that bf16 is just fp32 with the lower 16-bits of mantissa dropped;
# rather than hold 16-bit and 32-bit copies at once, we stash those lower
# 16 mantissa bits, and reconstruct the full 32-bit precision to update then
# resplit.
upper_bf16 = lambda w: (w.contiguous().view(torch.int32) >> 16).to(torch.int16).view(torch.bfloat16)
lower_uint16 = lambda w: (w.contiguous().view(torch.int32) ).to(torch.int16).view(torch.uint16)
# Set the seed so that every rank gets the same initialization -- no broadcast
# from a master rank needed.
torch.manual_seed(42)
torch.cuda.manual_seed(42)
m = Model()
# Written out one tensor per line, deliberately: the shape, the dtype, and
# therefore the memory cost of every weight and every piece of optimizer state
# is readable in one place, and the axis names say which dimension is sharded.
#
# Dtype scheme (hardcoded, stated per tensor below):
# - Matrix banks + lm_head: bf16 live + uint16 mantissa (fp32 master via the
# mantissa trick), fp32 gradients, fp32 moments.
# - Embedding tables (input_embeds, value_embeds): bf16 live + uint16 mantissa
# (fp32 master via the mantissa trick). This deviates from nanochat, which
# kept its embeddings plain bf16 and let AdamW update them in place -- we
# pair them with a mantissa so the one AdamW kernel serves everything,
# rather than carrying a second bf16-live variant. Overall our code
# ~matches the validation loss of the original.
# Gradients are bf16 -- these are the two biggest tensors in the model,
# fp32 grads would double their scatter traffic and (at world>1) comm bytes,
# and bf16 matches the autograd baseline's numerics (bf16 params -> bf16
# .grad). Everything else accumulates gradients in fp32.
# - Scalars (resid/x0 lambdas, smear, backout): fp32 live, no mantissa, same
# as they've always been. (Rounding them to bf16 was tried during the port
# and cost +0.016 val bpb, so they stay fp32.)
#
# Initialization values:
# input_embeds: normal, std=0.8
# lm_head: normal, std=0.001
# W_Q, W_K, W_V: uniform, bound=sqrt(3)/sqrt(d_model) -> std = 1/sqrt(d_model)
# W_O: zeros
# W_in: uniform, bound=0.4*sqrt(3)/sqrt(d_model) -> std = 0.4/sqrt(d_model)
# W_out: zeros
# value_embeds: uniform, bound=sqrt(3)/sqrt(d_model) (same as W_V)
# ve_gate: uniform in [0, 0.02] (slightly above neutral)
# resid_lambdas: 1.15 -> 1.05 linear decay over depth
# x0_lambdas: 0.20 -> 0.05 linear decay over depth
# smear_gate: zeros
# smear_lambda: zeros (smear disabled at init)
# backout_lambda: zeros (backout disabled at init)
# (Zeros for smear/backout is what nanochat's baselines actually trained
# with: it intended backout_lambda=0.2 and a kaiming smear_gate, but its
# meta-device init never ran those. Details at the Scalars block below.)
# Uniform init bound. Var(Uniform(-a, a)) = a^2/3, so std = a/sqrt(3): to hit
# a target std of 1/sqrt(d_model), the bound must be sqrt(3) times it.
matrix_init_s = (3 ** 0.5) * (cfg.d_model ** -0.5)
# ==== Input Embeddings ====
# bf16 live; draw in fp32 and let copy_ round -- drawing straight into bf16
# would quantize the distribution rather than the samples. The master upcast of
# a bf16 live is lossless, so the mantissa starts at zero.
# TODO - Leaving the zero-mantissa init for the moment (it matches the fwd-bwd
# reference), but we'll likely switch to keeping the draw's lower 16
# bits (the lm_head split pattern) once we have the chance to test.
m.input_embeds = bf16_empty(cfg.d_vocab, cfg.d_model)
m.input_embeds.copy_(fp32_empty(cfg.d_vocab, cfg.d_model).normal_(mean=0.0, std=0.8))
m.input_embeds.grad32 = bf16_zeros(cfg.d_vocab, cfg.d_model) # TODO - Change to `grad` since there's no colision?
m.input_embeds.mantissa = uint16_zeros(vocab_shard_size, cfg.d_model)
m.input_embeds.exp_avg = fp32_zeros(vocab_shard_size, cfg.d_model)
m.input_embeds.exp_avg_sq = fp32_zeros(vocab_shard_size, cfg.d_model)
# ==== Value Embeddings ====
# Same init std as W_V; same bf16-live / zero-mantissa path as input_embeds.
# AdamW state is shaped over the FLATTENED (ve_slot * vocab) row axis;
# optimizer_step passes matching 2-D views of the live bank and its grad.
# Flattening (vs a 3-D state mirroring the bank) is what lets ONE
# reduce-scatter/all-gather over dim-0 rows shard the whole bank evenly --
# per-slot vocab sharding on the 3-D layout would need a collective per VE
# slot. At world=1 a 3-D state would also work, but would have to reallocate
# the moment we go multi-GPU.
m.value_embeds = bf16_empty(cfg.num_ves, cfg.d_vocab, cfg.n_kv_heads * cfg.d_vo)
m.value_embeds.copy_(fp32_empty(cfg.num_ves, cfg.d_vocab, cfg.n_kv_heads * cfg.d_vo)
.uniform_(-matrix_init_s, matrix_init_s))
m.value_embeds.grad32 = bf16_zeros(cfg.num_ves, cfg.d_vocab, cfg.n_kv_heads * cfg.d_vo)
m.value_embeds.mantissa = uint16_zeros(ve_row_shard_size, cfg.n_kv_heads * cfg.d_vo)
m.value_embeds.exp_avg = fp32_zeros(ve_row_shard_size, cfg.n_kv_heads * cfg.d_vo)
m.value_embeds.exp_avg_sq = fp32_zeros(ve_row_shard_size, cfg.n_kv_heads * cfg.d_vo)
# ==== LM Head ====
# Drawn in fp32 and split -- unlike the embeddings, its mantissa is real from
# step zero.
lm_head_fp32 = fp32_empty(cfg.d_vocab, cfg.d_model).normal_(mean=0.0, std=0.001)
m.lm_head = upper_bf16(lm_head_fp32) # Live weights - bf16
m.lm_head.mantissa = lower_uint16(lm_head_fp32[vocab_shard_slice]) # Lower 16 bits for optimizer
del lm_head_fp32
m.lm_head.grad32 = fp32_zeros(cfg.d_vocab, cfg.d_model)
m.lm_head.exp_avg = fp32_zeros(vocab_shard_size, cfg.d_model)
m.lm_head.exp_avg_sq = fp32_zeros(vocab_shard_size, cfg.d_model)
# ==== Attention ====
# Parameter banks: the layer index is dim 0. Each slice uses F.linear's
# (out_features, in_features) convention and is consumed as `x @ w.mT`.
# Initialize in fp32 and split into bf16 live + uint16 mantissa.
W_Q_fp32 = fp32_empty(cfg.n_layers, cfg.n_q_heads * cfg.d_qk, cfg.d_model).uniform_(-matrix_init_s, matrix_init_s)
W_K_fp32 = fp32_empty(cfg.n_layers, cfg.n_kv_heads * cfg.d_qk, cfg.d_model).uniform_(-matrix_init_s, matrix_init_s)
W_V_fp32 = fp32_empty(cfg.n_layers, cfg.n_kv_heads * cfg.d_vo, cfg.d_model).uniform_(-matrix_init_s, matrix_init_s)
W_O_fp32 = fp32_zeros(cfg.n_layers, cfg.d_model, cfg.n_o_heads * cfg.d_vo) # projections start at zero
m.W_Q = upper_bf16(W_Q_fp32) # Live weights - bf16
m.W_K = upper_bf16(W_K_fp32)
m.W_V = upper_bf16(W_V_fp32)
m.W_O = upper_bf16(W_O_fp32)
# For the mantissa, we only need to hold our shard of the weights.
m.W_Q.mantissa = lower_uint16(W_Q_fp32[layer_shard_slice]) # Lower 16 bits for optimizer
m.W_K.mantissa = lower_uint16(W_K_fp32[layer_shard_slice])
m.W_V.mantissa = lower_uint16(W_V_fp32[layer_shard_slice])
m.W_O.mantissa = lower_uint16(W_O_fp32[layer_shard_slice])
del W_Q_fp32, W_K_fp32, W_V_fp32, W_O_fp32
# Gradients (full size -- the reduce-scatter source, never sharded)
m.W_Q.grad32 = fp32_zeros(cfg.n_layers, cfg.n_q_heads * cfg.d_qk, cfg.d_model)
m.W_K.grad32 = fp32_zeros(cfg.n_layers, cfg.n_kv_heads * cfg.d_qk, cfg.d_model)
m.W_V.grad32 = fp32_zeros(cfg.n_layers, cfg.n_kv_heads * cfg.d_vo, cfg.d_model)
m.W_O.grad32 = fp32_zeros(cfg.n_layers, cfg.d_model, cfg.n_o_heads * cfg.d_vo)
# First-momentum buffers for Muon (sharded)
m.W_Q.frst_mntm = fp32_zeros(layer_shard_size, cfg.n_q_heads * cfg.d_qk, cfg.d_model)
m.W_K.frst_mntm = fp32_zeros(layer_shard_size, cfg.n_kv_heads * cfg.d_qk, cfg.d_model)
m.W_V.frst_mntm = fp32_zeros(layer_shard_size, cfg.n_kv_heads * cfg.d_vo, cfg.d_model)
m.W_O.frst_mntm = fp32_zeros(layer_shard_size, cfg.d_model, cfg.n_o_heads * cfg.d_vo)
# Second momentum (NorMuon variance reduction) holds a running average of each
# neuron's mean-square update, so it is a vector (per layer) rather than a
# matrix mirroring the weights. (The neuron's rms is the square root of what's
# stored; the kernel applies it as an rsqrt.)
# NorMuon is a ~no-op for square matrices: polar express produces a
# ~orthonormal matrix, so the neuron norms are already ~uniform and there is
# nothing to normalize (confirmed with experiments). It only affects attention
# when the number of heads times the head size differs from d_model.
# The original code uses a heuristic to infer the neuron dimension by assuming
# that it is the smaller of the two. While typical, it's not certain. Instead,
# we specify it directly.
# Neurons can be identified directly by their interaction with the residual
# stream--they read from it and write to it and match it in length, so the
# mean-square is taken along the residual dimension.
# Note that the attention output projection consists of heads as well, and
# they are stored transposed relative to QKV, so we calculate the mean-square
# along dim -2.
m.W_Q.residual_dim = -1
m.W_K.residual_dim = -1
m.W_V.residual_dim = -1
m.W_O.residual_dim = -2
m.W_Q.scnd_mntm = fp32_zeros(layer_shard_size, cfg.n_q_heads * cfg.d_qk, 1)
m.W_K.scnd_mntm = fp32_zeros(layer_shard_size, cfg.n_kv_heads * cfg.d_qk, 1)
m.W_V.scnd_mntm = fp32_zeros(layer_shard_size, cfg.n_kv_heads * cfg.d_vo, 1)
m.W_O.scnd_mntm = fp32_zeros(layer_shard_size, 1, cfg.n_o_heads * cfg.d_vo)
# ==== MLPs ====
# For a transformer, 'MLP' is something of a misnomer. It's closer to a
# lookup table, containing pairs of vectors, both of length d_m.
# For a given pair (w_in, w_out), if the residual stream is positively
# aligned with w_in, then w_out is written back to it.
# But unlike a look up table, where a read-write operation is captured
# by a single row, here the model composes the operation across many
# vector pairs.
W_in_fp32 = fp32_empty(cfg.n_layers, cfg.d_mlp, cfg.d_model).uniform_(-matrix_init_s * 0.4, matrix_init_s * 0.4)
W_out_fp32 = fp32_zeros(cfg.n_layers, cfg.d_model, cfg.d_mlp) # projections start at zero
m.W_in = upper_bf16(W_in_fp32) # Live weights - bf16
m.W_out = upper_bf16(W_out_fp32)
m.W_in.mantissa = lower_uint16(W_in_fp32[layer_shard_slice]) # Lower 16 bits for optimizer
m.W_out.mantissa = lower_uint16(W_out_fp32[layer_shard_slice])
del W_in_fp32, W_out_fp32
# Gradients (full size)
m.W_in.grad32 = fp32_zeros(cfg.n_layers, cfg.d_mlp, cfg.d_model)
m.W_out.grad32 = fp32_zeros(cfg.n_layers, cfg.d_model, cfg.d_mlp)
# First-momentum buffers for Muon (sharded)
m.W_in.frst_mntm = fp32_zeros(layer_shard_size, cfg.d_mlp, cfg.d_model)
m.W_out.frst_mntm = fp32_zeros(layer_shard_size, cfg.d_model, cfg.d_mlp)
# Residual dimension: W_in rows read from the residual stream, W_out columns
# write to it.
m.W_in.residual_dim = -1
m.W_out.residual_dim = -2
m.W_in.scnd_mntm = fp32_zeros(layer_shard_size, cfg.d_mlp, 1)
m.W_out.scnd_mntm = fp32_zeros(layer_shard_size, 1, cfg.d_mlp)
# ==== VE Gates ====
# Muon, REPLICATED: tiny and ragged against world sizes, so every rank runs the
# full-size update rather than paying comm to shard a few thousand floats.
ve_gate_fp32 = fp32_empty(cfg.num_ves, cfg.n_kv_heads, cfg.d_ve_gate).uniform_(0.0, 0.02)
m.ve_gate = upper_bf16(ve_gate_fp32)
m.ve_gate.mantissa = lower_uint16(ve_gate_fp32) # replicated: full-size mantissa
del ve_gate_fp32
m.ve_gate.grad32 = fp32_zeros(cfg.num_ves, cfg.n_kv_heads, cfg.d_ve_gate)
m.ve_gate.frst_mntm = fp32_zeros(cfg.num_ves, cfg.n_kv_heads, cfg.d_ve_gate)
m.ve_gate.residual_dim = -1 # gate rows read a d_ve_gate slice of the residual stream
m.ve_gate.scnd_mntm = fp32_zeros(cfg.num_ves, cfg.n_kv_heads, 1)
# ==== Scalars ====
# fp32-LIVE with no mantissa pair (see the dtype scheme note above). AdamW,
# replicated.
# These serve separate purposes:
# - resid_lambdas: Directly scales the residual stream at the start of each layer.
# - x0_lambdas: How strongly the input embedding is added to the residual stream.
# Per-layer scalars: linear decay over depth. Stronger residual and more
# input-embedding blending at early layers, both tapering with depth.
m.resid_lambdas = torch.linspace(1.15, 1.05, cfg.n_layers, dtype=torch.float32, device=device)
m.x0_lambdas = torch.linspace(0.20, 0.05, cfg.n_layers, dtype=torch.float32, device=device)
# Smear/backout start disabled, zeros everywhere.
# Note: nanochat pre-flattening had a bug here--it intended backout_lambda=0.2
# and a kaiming smear_gate, but under meta-device init those never executed and
# to_empty() left zeroed storage. Zeros is what every tuned baseline actually
# trained with, so now it's explicit rather than luck.
m.smear_gate = fp32_zeros(1, cfg.d_smr_gate)
m.smear_lambda = fp32_zeros(1)
m.backout_lambda = fp32_zeros(1)
m.resid_lambdas.grad32 = fp32_zeros(cfg.n_layers)
m.x0_lambdas.grad32 = fp32_zeros(cfg.n_layers)
m.smear_gate.grad32 = fp32_zeros(1, cfg.d_smr_gate)
m.smear_lambda.grad32 = fp32_zeros(1)
m.backout_lambda.grad32 = fp32_zeros(1)
m.resid_lambdas.exp_avg = fp32_zeros(cfg.n_layers)
m.resid_lambdas.exp_avg_sq = fp32_zeros(cfg.n_layers)
m.x0_lambdas.exp_avg = fp32_zeros(cfg.n_layers)
m.x0_lambdas.exp_avg_sq = fp32_zeros(cfg.n_layers)
m.smear_gate.exp_avg = fp32_zeros(1, cfg.d_smr_gate)
m.smear_gate.exp_avg_sq = fp32_zeros(1, cfg.d_smr_gate)
m.smear_lambda.exp_avg = fp32_zeros(1)
m.smear_lambda.exp_avg_sq = fp32_zeros(1)
m.backout_lambda.exp_avg = fp32_zeros(1)
m.backout_lambda.exp_avg_sq = fp32_zeros(1)
# ==== Rotary Cache ====
# Without an nn.Module these are just attributes on m -- register_buffer only
# existed for state_dict/.to() plumbing we no longer have (and these were
# persistent=False anyway). With varlen training the whole micro-batch is one
# packed sequence, so the cache spans the largest T any forward sees: the
# training micro-batch (val micro-batches match it) or the CORE/chat eval
# packing buffer, whichever is bigger. The assert in the forward bodies
# catches it if we ever exceed.
rotary_seq_len = max(cfg.micro_batch_tokens, cfg.eval_buffer_tokens)
channel_range = torch.arange(0, cfg.d_qk, 2, dtype=torch.float32, device=device) # stride the channels
inv_freq = 1.0 / (100000 ** (channel_range / cfg.d_qk))
t_pos = torch.arange(rotary_seq_len, dtype=torch.float32, device=device) # stride the time steps
freqs = torch.outer(t_pos, inv_freq) # rotation frequency at each (time, channel) pair
m.cos = freqs.cos().to(torch.bfloat16)[None, :, None, :] # add batch and head dims
m.sin = freqs.sin().to(torch.bfloat16)[None, :, None, :] # for later broadcasting
del channel_range, inv_freq, t_pos, freqs
# ==== Bank Gradient Slice Views ====
# The 3-D banks get `grad32_slices`: per-slice VIEWS built OUTSIDE any compiled
# graph. The forward/backward bodies accumulate through these, never through
# `grad32[i]` -- an in-graph bank slice functionalizes into a whole-bank
# select_scatter copy (10-20x the cost of the slice add at these bank sizes),
# while a view created out of graph arrives as an input and mutates genuinely
# in place.
m.W_Q.grad32_slices = list(m.W_Q.grad32.unbind(0))
m.W_K.grad32_slices = list(m.W_K.grad32.unbind(0))
m.W_V.grad32_slices = list(m.W_V.grad32.unbind(0))
m.W_O.grad32_slices = list(m.W_O.grad32.unbind(0))
m.W_in.grad32_slices = list(m.W_in.grad32.unbind(0))
m.W_out.grad32_slices = list(m.W_out.grad32.unbind(0))
m.ve_gate.grad32_slices = list(m.ve_gate.grad32.unbind(0))
m.value_embeds.grad32_slices = list(m.value_embeds.grad32.unbind(0))
# (Grad zeroing happens as a plain loop at the training-loop call site --
# every .grad32 is zeroed after each optimizer_step, since gradients
# accumulate across a step's micro-batches AND Muon's nesterov lerp mutates
# grad32 in place.)
# --------------------------------------------------------------------------------
# § Schedules
# --------------------------------------------------------------------------------
# A run's optimizer is defined up front: every learning rate, beta and weight
# decay for every step is computed here, before training starts, into per-step
# tables of *update coefficients* -- the numbers the fused kernels actually
# multiply by. The optimizer then holds no hyperparameters of its own and the
# training loop has nothing to set per step; the kernels just gather row
# `t_step` of each table. Folding all the way down to coefficients buys:
# - The bias corrections leave the kernel (betas are per-role constants, so
# the closed `1 - beta^t` form is exact).
# - Nothing about the schedule is left for the loop to do per step. Tables are
# device-resident and the step counter is a device tensor, so a step involves
# the host for nothing at all.
class AdamWTabs(NamedTuple):
"""What an AdamW step multiplies by, one (N,) table per field. eps is never
scheduled, so it rides as a plain kernel argument instead of a table."""
wd_mul: Tensor # 1 - lr*wd decoupled weight decay
one_minus_beta1: Tensor # 1 - beta1 exp_avg lerp weight
one_minus_beta2: Tensor # 1 - beta2 exp_avg_sq lerp weight
rsqrt_bias2: Tensor # 1/sqrt(bias2) second-moment bias correction
step_size: Tensor # lr / bias1 lr schedule x first-moment bias correction
class MuonCoeffs(NamedTuple):
"""What a Muon step multiplies by. Muon's second moment is self-normalizing
(the v_norm/v_norm_new rescale), so it needs no bias correction."""
momentum: Tensor # nesterov momentum
one_minus_momentum: Tensor # 1 - momentum frst_mntm lerp weight
one_minus_beta2: Tensor # 1 - beta2 variance-reduction lerp weight
lr: Tensor # lr (the per-bank aspect scale arrives separately, via lr_mul)
lr_wd: Tensor # lr * weight_decay cautious decay
def build_schedules(num_iterations, batch_lr_scale=1.0, weight_decay=0.28,
warmup_steps=40, warmdown_ratio=0.65, final_lr_frac=0.05):
"""Named table sets with the tuned nanochat base_train hyperparameters,
written out flat. Baked assumptions (a Ramp class used to support more):
exactly three shaped schedules exist -- the shared LR multiplier, Muon
momentum, and Muon weight decay; every Adam beta is a per-role CONSTANT;
windows are warmup_steps + round(warmdown_ratio * N). Verified
bitwise-identical to the Ramp implementation it replaced
(sched_parity_test.py in the session folder).
`weight_decay` arrives already batch/horizon-scaled. Returns a namespace:
.matrix (MuonCoeffs) + one AdamWTabs per AdamW role, .adamw_eps, and
.num_steps. The trainer binds the result to the global `sched`."""
N = num_iterations
C = round(warmdown_ratio * N) # LR warmdown length
assert warmup_steps + C <= N, f"warmup ({warmup_steps}) + warmdown ({C}) exceed the run ({N})"
i = np.arange(N, dtype=np.float64)
cool = slice(N - C + 1, N) # the hold covers i <= N - C
f = (N - i[cool]) / C # ~1 -> ~0 across the warmdown
# The one LR shape for the whole run: linear warmup from 0 (reaching the
# peak on the warmup window's last step), hold at 1, linear warmdown to
# final_lr_frac (arriving one step past the run's end -- nanochat's
# convention). Each role scales it to its own peak below.
lrm = np.ones(N)
lrm[:warmup_steps] = (i[:warmup_steps] + 1.0) / warmup_steps
lrm[cool] = final_lr_frac + (1.0 - final_lr_frac) * f
# Muon momentum: 0.85 -> 0.97 over 400 steps (the clamp only lets short
# smoke/debug runs build a valid schedule; identical for N >= ~1150),
# hold, then cool to 0.90 across the LR warmdown.
mW = min(400, int(N * (1 - warmdown_ratio)))
momentum = np.full(N, 0.97)
momentum[:mW] = 0.85 + (0.97 - 0.85) * (i[:mW] + 1.0) / mW
momentum[cool] = 0.90 + (0.97 - 0.90) * f
# Muon weight decay: half-cosine from the peak to zero over the whole run
# (step 0 sits at the peak; the decay begins at step 1).
muon_wd = np.empty(N)
muon_wd[0] = weight_decay
fw = (N - i[1:]) / N
muon_wd[1:] = weight_decay * (0.5 * (1.0 + np.cos(math.pi * (1.0 - fw))))
# Numpy arrays -> fp32 device tables: a step reads its coefficients with
# an on-device gather, never a host-to-device copy.
dev = lambda a: torch.tensor(a, dtype=torch.float32, device=device)
t1 = np.arange(1, N + 1, dtype=np.float64)
# Fold one AdamW role's schedule down to the kernel's update coefficients.
# The folding is ONE policy shared by every role (repeating it 6x would
# obscure edits); the per-role peaks/betas/wd stay visible at the call
# sites below.
def adamw(peak, beta1, beta2, wd):
lr = lrm * peak
return AdamWTabs(
wd_mul = dev(1.0 - lr * wd),
one_minus_beta1 = dev(np.full(N, 1.0 - beta1)),
one_minus_beta2 = dev(np.full(N, 1.0 - beta2)),
rsqrt_bias2 = dev(1.0 / ((1.0 - beta2 ** t1) ** 0.5)),
step_size = dev(lr / (1.0 - beta1 ** t1)),
)
# Muon's coefficients fold directly from the three shaped schedules.
# Canonical lr: NO per-bank aspect fold (see § Optimizer Step).
matrix_lr = lrm * (0.02 * batch_lr_scale)
matrix = MuonCoeffs(
momentum = dev(momentum),
one_minus_momentum = dev(1.0 - momentum),
one_minus_beta2 = dev(np.full(N, 1.0 - 0.9)), # variance-reduction beta2 = 0.9
lr = dev(matrix_lr),
lr_wd = dev(matrix_lr * muon_wd),
)
# Per-role peak LRs (tuned values). The AdamW peaks were tuned at d12's
# width, so they carry the 1/sqrt(width ratio) correction to d24.
adamw_lr_scale = batch_lr_scale * (cfg.d_model / 768) ** -0.5
return SimpleNamespace(
matrix = matrix,
lm_head = adamw(0.008 * adamw_lr_scale, 0.8, 0.96, 0.01),
input_embeds = adamw(0.3 * adamw_lr_scale, 0.8, 0.995, 0.001),
value_embeds = adamw(0.3 * adamw_lr_scale * 0.5, 0.8, 0.995, 0.01),
resid = adamw(0.5 * batch_lr_scale * 0.01, 0.8, 0.95, 0.05),
x0 = adamw(0.5 * batch_lr_scale, 0.96, 0.95, 0.0),
smear = adamw(0.2, 0.8, 0.95, 0.0),
adamw_eps = 1e-10,
lrm_table = lrm, # host-side copy, for logging only
num_steps = N,
batch_lr_scale = batch_lr_scale, # echoed into the wandb config
weight_decay = weight_decay,
)
# --------------------------------------------------------------------------------
# § Optimizer Code
# --------------------------------------------------------------------------------
# --------------------------------------------------------------------------------
# Mantissa Trick
# Masters use the mantissa trick (Larry Dial via modded-nanogpt train_gpt.py):
# the fp32 master's bit pattern is (live_bf16_bits << 16) | mantissa_uint16.
# Update math runs in fp32 on the reconstructed master; the split back is a
# TRUNCATION (load-bearing: round-to-nearest could carry into the top bits and
# break the lossless live/mantissa pairing).
#
# The bit arithmetic runs in int32 (CUDA has no uint32 shifts as of torch 2.9);
# int32's truncating .to(int16) and the <<16 discard of sign-extension bits
# make it equivalent. Mantissa tensors are STORED uint16, viewed int16 for the
# math.
def fp32_master(live: Tensor, mantissa: Tensor) -> Tensor:
"""Reconstruct the fp32 master from bf16 live bits + stashed mantissa."""
bits = (live.view(torch.int16).to(torch.int32) << 16) | \
(mantissa.view(torch.int16).to(torch.int32) & 0xFFFF)
return bits.view(torch.float32)
def writeback_master(master: Tensor, live: Tensor, mantissa: Tensor) -> None:
"""Truncation split of the updated master back into live + mantissa."""
bits = master.view(torch.int32)
live.view(torch.int16).copy_((bits >> 16).to(torch.int16))
mantissa.view(torch.int16).copy_(bits.to(torch.int16))
# -----------------------------------------------------------------------------
# Fused update kernels. The schedule row is gathered ON DEVICE by `t` -- no
# host involvement per step.
# We use the first five, remainder are just for completeness.
polar_express_coeffs = [
(8.156554524902461, -22.48329292557795, 15.878769915207462),
(4.042929935166739, -2.808917465908714, 0.5000178451051316),
(3.8916678022926607, -2.772484153217685, 0.5060648178503393),
(3.285753657755655, -2.3681294933425376, 0.46449024233003106),
(2.3465413258596377, -1.7097828382687081, 0.42323551169305323),
]
@torch.compile(dynamic=False, fullgraph=True)
def adamw_step_fused_fp32(
p: Tensor, # fp32 param, updated IN PLACE (live == master)
grad: Tensor,
exp_avg: Tensor,
exp_avg_sq: Tensor,
c: AdamWTabs,
t: Tensor, # (1,) int64 device tensor - the schedule row to read
eps: float,
) -> None:
"""AdamW for the fp32-LIVE scalar params (resid/x0 lambdas, smear, backout
-- ~30 floats). They are exempt from the bf16-live/mantissa scheme: see the
dtype scheme note in § Model Initialization."""
grad = grad.to(exp_avg.dtype)
p.mul_(c.wd_mul[t])
exp_avg.lerp_(grad, c.one_minus_beta1[t])
exp_avg_sq.lerp_(grad.square(), c.one_minus_beta2[t])
denom = exp_avg_sq.sqrt() * c.rsqrt_bias2[t] + eps
p.sub_(c.step_size[t] * (exp_avg / denom))
@torch.compile(dynamic=False, fullgraph=True)
def adamw_step_fused(
live: Tensor, # bf16 live shard
mantissa: Tensor, # uint16, same shape
grad: Tensor, # gradient shard (fp32, or bf16 for the embeddings)
exp_avg: Tensor, # fp32 first moment
exp_avg_sq: Tensor, # fp32 second moment
c: AdamWTabs, # per-step coefficient tables, device-resident
t: Tensor, # (1,) int64 device tensor - the schedule row to read
eps: float,
) -> None:
"""Fused AdamW step on the reconstructed master."""
p = fp32_master(live, mantissa)
grad = grad.to(exp_avg.dtype) # embeddings hand in bf16 grads; moment math stays fp32
p.mul_(c.wd_mul[t])
exp_avg.lerp_(grad, c.one_minus_beta1[t])
exp_avg_sq.lerp_(grad.square(), c.one_minus_beta2[t])
denom = exp_avg_sq.sqrt() * c.rsqrt_bias2[t] + eps
p.sub_(c.step_size[t] * (exp_avg / denom))
writeback_master(p, live, mantissa)
# The update kernels take explicit per-tensor arguments rather than the model
# object, twice over: (1) at world>1 the SAME kernels run on shard views
# (p[layer_shard_slice] with the shard-size state) rather than on m.X -- an
# object-reading kernel would need a different body per world size; (2) under
# fullgraph compile,
# attribute access on an ad-hoc Python object turns into dynamo guards on
# object identity/attributes -- fragile and recompile-prone next to plain
# tensor arguments.
@torch.compile(dynamic=False, fullgraph=True)
def muon_step_fused(
grad: Tensor, # (K, out, in) fp32 gradient shard -- MUTATED (nesterov lerp)
live: Tensor, # (K, out, in) bf16 live shard
mantissa: Tensor, # (K, out, in) uint16
frst_mntm: Tensor, # (K, out, in) fp32
scnd_mntm: Tensor, # (K, out, 1) or (K, 1, in) fp32 - factored second moment
c: MuonCoeffs, # per-step coefficient tables, device-resident (UNfolded lr)
t: Tensor, # (1,) int64 device tensor - the schedule row to read
ns_steps: int, # 5 - number of Polar Express iterations
residual_dim: int, # -1 or -2 - residual-facing axis; per-neuron mean-square is taken along it
lr_mul: Tensor, # (K, 1, 1) fp32 per-slice LR multiplier (aspect scale today)
wd_mul: Tensor, # (K, 1, 1) fp32 per-slice WD multiplier
) -> None:
"""Fused Muon step: momentum -> polar_express -> variance_reduction ->
cautious update on the reconstructed master. The sqrt(fan_out/fan_in)
aspect scale is NOT in `c` -- it arrives through lr_mul/wd_mul, per slice,
so the one coefficient table stays valid for every bank."""
dtype = grad.dtype
# Nesterov momentum
frst_mntm.lerp_(grad, c.one_minus_momentum[t].to(dtype))
g = grad.lerp_(frst_mntm, c.momentum[t].to(dtype))
# Polar express (orthogonalization)
X = g.bfloat16()
X = X / (X.norm(dim=(-2, -1), keepdim=True) * 1.01 + 1e-6)
if g.size(-2) > g.size(-1): # Tall matrix
for a, b, c_ns in polar_express_coeffs[:ns_steps]:
A = X.mT @ X
B = b * A + c_ns * (A @ A)
X = a * X + X @ B
else: # Wide matrix (original math)
for a, b, c_ns in polar_express_coeffs[:ns_steps]:
A = X @ X.mT
B = b * A + c_ns * (A @ A)
X = a * X + B @ X
g = X
# Variance reduction (NorMuon). The lerp weight stays fp32.
v_mean = g.float().square().mean(dim=residual_dim, keepdim=True)
residual_dim_size = g.size(residual_dim)
v_norm_sq = v_mean.sum(dim=(-2, -1), keepdim=True) * residual_dim_size
v_norm = v_norm_sq.sqrt()
scnd_mntm.lerp_(v_mean.to(dtype=scnd_mntm.dtype),
c.one_minus_beta2[t].to(scnd_mntm.dtype))
step_size = scnd_mntm.clamp_min(1e-10).rsqrt()
scaled_sq_sum = (v_mean * residual_dim_size) * step_size.float().square()
v_norm_new = scaled_sq_sum.sum(dim=(-2, -1), keepdim=True).sqrt()
final_scale = step_size * (v_norm / v_norm_new.clamp_min(1e-10))
g = g * final_scale.to(g.dtype)
# Cautious weight decay + master update + truncation split back to live
p = fp32_master(live, mantissa)
mask = (g * p) >= 0
lr = (c.lr[t] * lr_mul).to(g.dtype)
lr_wd = (c.lr_wd[t] * wd_mul).to(g.dtype)
p.sub_(lr * g + lr_wd * p * mask)
writeback_master(p, live, mantissa)
# --------------------------------------------------------------------------------
# § Model Code (Forward/Backward)
# --------------------------------------------------------------------------------
# Handwritten training step: explicit forward + backward (no autograd),
# accumulating into the fp32/bf16 `.grad32` buffers.
#
# Design notes:
# - Attention runs through the raw FA3 ops above, stashing out + LSE.
# - rms_norms: we stash the norm OUTPUT plus the per-vector 1/rms `r`. In
# output space the backward is dx = r*(dy - y*mean(y*dy)) for ANY eps, so the
# pre-norm input is never needed. Cheap norms (the MLP-side xm) are
# recomputed from the stashed pre-norm x1 instead of stashed.
# - Weight-grad matmuls run in bf16, then accumulate upcast into grad32 -- the
# same numerics autograd produces for a bf16 matmul.
# - loss_scale (1/grad_accum_steps) replaces the loss division of an autograd
# loop; the returned loss is the plain (unscaled) mean CE for logging.
# Cast shorthands for the bodies below: the fp32 scalars/gates need explicit
# bf16 casts at their use sites (see forward_backward's docstring), and the
# scalar-parameter grad sums accumulate in fp32.
bf16 = lambda x: x.to(torch.bfloat16)
sum32 = lambda x: x.sum(dtype=torch.float32)
# -----------------------------------------------------------------------------
# rms_norm forward/backward in output space
# TODO - Inline at call site. And can one not derive the other?
def _rms_fwd(x):
"""rms_norm over the last dim plus the per-vector 1/rms its backward
needs, sharing one mean-square. r is fp32 with eps = 2^-23 (fp32 machine
eps -- the same number compiled F.rms_norm's decomposition uses); y is
x * r cast back to bf16. Verified bitwise-identical to the F.rms_norm
form under torch.compile, and the same speed (bench_rms.log; eager ATen
differs in last-ulp on ~6/1M elements, but every call site is compiled)."""
r = (x.float().square().mean(dim=-1, keepdim=True) + 2.0 ** -23).rsqrt()
y = bf16(x.float() * r)
return y, r
# TODO - Inline.
def _rms_bwd(dy, y, r):
"""dx = r*(dy - y*mean(y*dy)): exact for any eps because r is the forward's
actual 1/rms and y the actual output (substitute x = y/r in the usual
form). Math in fp32, result back to bf16."""
yf, dyf = y.float(), dy.float()
dx = r * (dyf - yf * (yf * dyf).mean(dim=-1, keepdim=True))
return bf16(dx)
# TODO - Inline.
def _rms_bwd_scaled(dy, ys, r, s):
"""Backward through ys = s * rms_norm(x), given the SCALED output ys --
which is exactly what the attention kernel consumed, so it stashes directly
with no recompute pass. Substituting y = ys/s into _rms_bwd's form:
dx = r*(s*dy - ys*mean(ys*dy)/s). Exact algebra."""
yf, dyf = ys.float(), dy.float()
dx = r * (s * dyf - yf * ((yf * dyf).mean(dim=-1, keepdim=True) / s))
return bf16(dx)
# -----------------------------------------------------------------------------
# forward_backward
@torch.no_grad()
def forward_backward(idx, targets, cu_seqlens, loss_scale=1.0):
"""One micro-batch: forward, stash, explicit backward into `.grad32`.
Returns the detached mean CE loss (unscaled; grads carry loss_scale).
Wrap in torch.compile -- the CE block below is written for inductor's
fusion; run eager it materializes full (T, d_vocab) fp32 temporaries.
Activations are bf16 throughout. The live weights are already bf16, so no
per-use casts; the fp32 scalars need care: indexing a 1-D fp32 bank gives a
0-dim tensor, which does NOT promote a bf16 tensor (resid/x0 lambdas ride
as-is), but the (1,)-shaped smear/backout scalars and the smear_gate matrix
WOULD promote to fp32, so those are cast explicitly."""
assert idx.ndim == 1
T = idx.size(0)
nl = cfg.n_layers
nh, nkv = cfg.n_q_heads, cfg.n_kv_heads
dqk, dvo = cfg.d_qk, cfg.d_vo
half = dqk // 2
gch = cfg.d_ve_gate
assert T > 1, "Training forward pass should have T > 1"
assert T <= m.cos.size(1), f"Sequence length grew beyond the rotary embeddings cache: {T} > {m.cos.size(1)}"
cos, sin = m.cos[0, :T], m.sin[0, :T] # (T, 1, half)
# ==== forward half (mirrors forward() -- keep the two visibly line-parallel) ====
x = F.embedding(idx, m.input_embeds) # bf16
xe, r_e = _rms_fwd(x) # post-norm embedding, pre-smear
# Smear: mix the previous token's embedding into the current position.
gate = bf16(m.smear_lambda) * torch.sigmoid(
xe[1:, :cfg.d_smr_gate] @ bf16(m.smear_gate).mT)
x = torch.cat([xe[:1], xe[1:] + gate * xe[:-1]], dim=0)
x0 = x
backout_layer = nl // 2
x_backout = None
stash = []
for i in range(nl):
x_in = x
b = m.resid_lambdas[i] * x_in + m.x0_lambdas[i] * x0
xn, r_xn = _rms_fwd(b)
q = (xn @ m.W_Q[i].mT).view(T, nh, dqk)
k = (xn @ m.W_K[i].mT).view(T, nkv, dqk)
v = (xn @ m.W_V[i].mT).view(T, nkv, dvo)
j = cfg.ve_index[i]
if j >= 0:
ve = F.embedding(idx, m.value_embeds[j]).view(T, nkv, dvo)
g = 3 * torch.sigmoid(xn[..., :gch] @ m.ve_gate[j].mT)
v = v + g.unsqueeze(-1) * ve # ve/g recomputed in backward, not stashed
q1, q2 = q[..., :half], q[..., half:]
k1, k2 = k[..., :half], k[..., half:]
q = torch.cat([q1 * cos + q2 * sin, q1 * (-sin) + q2 * cos], dim=-1)
k = torch.cat([k1 * cos + k2 * sin, k1 * (-sin) + k2 * cos], dim=-1)
qn, r_q = _rms_fwd(q)
kn, r_k = _rms_fwd(k)
qf = qn * 1.2 # stash the SCALED q/k (the kernel's inputs);
kf = kn * 1.2 # backward folds the 1.2 via _rms_bwd_scaled
y, lse = flash_attn_varlen_fwd_lse(qf, kf, v, cu_seqlens, cfg.seq_len, cfg.window_sizes[i])
y = y.contiguous()
x1 = b + y.view(T, -1) @ m.W_O[i].mT
xm, _ = _rms_fwd(x1) # xm recomputed in backward from stashed x1
a = F.relu(xm @ m.W_in[i].mT)
x = x1 + a.square() @ m.W_out[i].mT
if i == backout_layer:
x_backout = x
stash.append(dict(x_in=x_in, xn=xn, r_xn=r_xn, qf=qf, kf=kf, r_q=r_q, r_k=r_k,
v=v, y=y, lse=lse, x1=x1, a=a))
x_pre = x - bf16(m.backout_lambda) * x_backout
xf, r_f = _rms_fwd(x_pre)
# lm_head + softcap + CE loss + dlogits, written for inductor's fusion:
# tcap is an explicit CSE target (materialize once, no tanh recompute in
# the dz pass), and the onehot is a broadcast compare (a scatter_add here
# forces an extra full pass over the buffer). Vocab is unpadded by
# construction, so there is no [:V] cropping anywhere. No pad/ignore
# machinery either: every target is a real token by construction (the
# loader packs whole documents; at a doc seam the target is the next
# doc's BOS), so the mean runs over all T positions and the dz scale is
# the compile-time constant loss_scale/T rather than a device n_valid.
softcap = 15.0
logits = xf @ m.lm_head.mT # (T, d_vocab) bf16
tcap = torch.tanh(logits.float() / softcap)
cap = softcap * tcap
tgt = targets.unsqueeze(1)
cap_y = cap.gather(1, tgt).squeeze(1)
cmax = cap.amax(dim=1, keepdim=True)
e = (cap - cmax).exp()
ssum = e.sum(dim=1, keepdim=True)
lse_ce = (ssum.log() + cmax).squeeze(1)
loss = (lse_ce - cap_y).mean()
onehot = torch.arange(cfg.d_vocab, device=targets.device).unsqueeze(0) == tgt
dz = bf16((e / ssum - onehot.float()) * (1.0 - tcap * tcap) * (loss_scale / T))
del logits
m.lm_head.grad32.add_((dz.mT @ xf).float())
dxf = dz @ m.lm_head
del dz
# ==== backward half ====
# Bank wgrads add directly into grad32_slices views; only the per-layer
# scalar sums are collected and landed stacked at the end.
g_resid = []; g_x0 = []
d_pre = _rms_bwd(dxf, xf, r_f)
m.backout_lambda.grad32.add_(-sum32(d_pre * x_backout))
d_stream = d_pre # grad wrt layer nl-1's output
d_x0 = torch.zeros_like(x0)
for i in reversed(range(nl)):
st = stash[i]
if i == backout_layer:
# TRAP: x_backout gets an EXTRA contribution when the sweep passes nl//2
d_stream = d_stream - bf16(m.backout_lambda) * d_pre
# --- MLP backward (relu^2: dh = 2*a*du, self-masking since a = relu(h)) ---
x1, a = st["x1"], st["a"]
d_u = d_stream @ m.W_out[i]
m.W_out.grad32_slices[i].add_(d_stream.mT @ a.square())
d_h = 2.0 * a * d_u
xm, r_xm = _rms_fwd(x1) # cheap recompute (bitwise: same input)
m.W_in.grad32_slices[i].add_(d_h.mT @ xm)
d_xm = d_h @ m.W_in[i]
d_x1 = d_stream + _rms_bwd(d_xm, xm, r_xm)
# --- attention backward ---
xn, y = st["xn"], st["y"]
m.W_O.grad32_slices[i].add_(d_x1.mT @ y.view(T, -1))
d_y = (d_x1 @ m.W_O[i]).view(T, nh, dvo)
dqf, dkf, dv = flash_attn_varlen_bwd(
d_y, st["qf"], st["kf"], st["v"], y, st["lse"], cu_seqlens, cfg.seq_len,
cfg.window_sizes[i])
# per-(token, head) norm backward with the 1.2 scale folded in
d_qr = _rms_bwd_scaled(dqf, st["qf"], st["r_q"], 1.2)
d_kr = _rms_bwd_scaled(dkf, st["kf"], st["r_k"], 1.2)
# rotary backward = rotation by -theta (transpose of the forward rotation)
dq1, dq2 = d_qr[..., :half], d_qr[..., half:]
d_q0 = torch.cat([dq1 * cos - dq2 * sin, dq1 * sin + dq2 * cos], dim=-1)
dk1, dk2 = d_kr[..., :half], d_kr[..., half:]
d_k0 = torch.cat([dk1 * cos - dk2 * sin, dk1 * sin + dk2 * cos], dim=-1)
# --- VE gate backward (ve/g recomputed) ---
j = cfg.ve_index[i]
d_xn_ve = None
if j >= 0:
ve = F.embedding(idx, m.value_embeds[j]).view(T, nkv, dvo)
sg = torch.sigmoid(xn[..., :gch] @ m.ve_gate[j].mT)
d_g = (dv * ve).sum(dim=-1) # (T, n_kv_heads)
d_zg = d_g * (3 * sg * (1 - sg))
m.ve_gate.grad32_slices[j].add_(d_zg.mT @ xn[..., :gch])
d_ve = (dv * (3 * sg).unsqueeze(-1)).reshape(T, nkv * dvo)
# embedding_dense_backward (autograd's own lowering) beats raw
# index_add_ atomics ~2x at these shapes -- see the GH200 trace hunt
m.value_embeds.grad32_slices[j].add_(
torch.ops.aten.embedding_dense_backward(d_ve, idx, cfg.d_vocab, -1, False))
d_xn_ve = d_zg @ m.ve_gate[j]
# dv passes through the VE add unchanged: v = v0 + g*ve
d_q0 = d_q0.view(T, nh * dqk)
d_k0 = d_k0.view(T, nkv * dqk)
d_v0 = dv.reshape(T, nkv * dvo)
m.W_Q.grad32_slices[i].add_(d_q0.mT @ xn)
m.W_K.grad32_slices[i].add_(d_k0.mT @ xn)
m.W_V.grad32_slices[i].add_(d_v0.mT @ xn)
d_xn = d_q0 @ m.W_Q[i] + d_k0 @ m.W_K[i] + d_v0 @ m.W_V[i]
if d_xn_ve is not None:
d_xn[:, :gch] += d_xn_ve
d_b = d_x1 + _rms_bwd(d_xn, xn, st["r_xn"])
# --- blend backward: b = resid_lambdas[i]*x_in + x0_lambdas[i]*x0 ---
g_resid.append(sum32(d_b * st["x_in"]))
g_x0.append(sum32(d_b * x0))
d_x0 = d_x0 + m.x0_lambdas[i] * d_b # TRAP: x0 feeds every layer, accumulate
d_stream = m.resid_lambdas[i] * d_b
stash[i] = None # free this layer's stash as we go
# Land the per-layer resid/x0 scalar sums (collected in REVERSED layer
# order) as one stacked add each.
m.resid_lambdas.grad32.add_(torch.stack(g_resid[::-1]))
m.x0_lambdas.grad32.add_(torch.stack(g_x0[::-1]))
# d_stream is now the grad through layer 0's input, which IS x0 (same tensor)
d_xs = d_x0 + d_stream # grad wrt the smeared embedding
# --- smear backward: xs = cat([xe[:1], xe[1:] + gate*xe[:-1]]) ---
sg = torch.sigmoid(xe[1:, :cfg.d_smr_gate] @ bf16(m.smear_gate).mT) # (T-1, 1), recomputed
gate = bf16(m.smear_lambda) * sg
d_xe = d_xs.clone()
d_xe[:-1] += gate * d_xs[1:] # TRAP: shifted scatter -- p's grad reaches p-1
d_gate = (d_xs[1:] * xe[:-1]).sum(dim=-1, keepdim=True) # (T-1, 1)
m.smear_lambda.grad32.add_(sum32(d_gate * sg))
d_zs = d_gate * bf16(m.smear_lambda) * sg * (1 - sg)
m.smear_gate.grad32.add_((d_zs.mT @ xe[1:, :cfg.d_smr_gate]).float())
d_xe[1:, :cfg.d_smr_gate] += d_zs @ bf16(m.smear_gate)
# --- embedding norm + token embedding scatter ---
d_emb = _rms_bwd(d_xe, xe, r_e)
m.input_embeds.grad32.add_(
torch.ops.aten.embedding_dense_backward(d_emb, idx, cfg.d_vocab, -1, False))
return loss
# --------------------------------------------------------------------------------
# § Forward-Only
# --------------------------------------------------------------------------------
# Compiled by the trainer: § Main Loop rebinds this name through torch.compile
# (one specialization per shape/targets combination -- val loss and CORE logits).
@torch.no_grad()
def forward(idx, cu_seqlens, targets=None, loss_reduction='mean'):
"""Scoring forward for validation loss and CORE eval: one packed 1D
sequence of documents with per-document attention isolation via varlen
flash attention. idx/targets are (T,) and activations stay (T, ...)
throughout -- the layout the varlen kernel wants. Returns the loss if
targets are given, else the (softcapped, fp32) logits (T, d_vocab).
Mirrors forward_backward's forward half line for line -- keep them that
way; diff them when either changes."""
assert idx.ndim == 1
T = idx.size(0)
D = cfg.d_model
half = cfg.d_qk // 2
assert T > 1, "Scoring forward pass should have T > 1 (smear needs a previous token)"
assert T <= m.cos.size(1), f"Sequence length grew beyond the rotary embeddings cache: {T} > {m.cos.size(1)}"
cos, sin = m.cos[0, :T], m.sin[0, :T] # (T, 1, half)
# Embed the tokens
x = F.embedding(idx, m.input_embeds) # bf16
x = F.rms_norm(x, (D,))
# Smear: mix the previous token's embedding into the current position.
gate = bf16(m.smear_lambda) * torch.sigmoid(
x[1:, :cfg.d_smr_gate] @ bf16(m.smear_gate).mT)
x = torch.cat([x[:1], x[1:] + gate * x[:-1]], dim=0)
# Forward the trunk of the Transformer
x0 = x
backout_layer = cfg.n_layers // 2
x_backout = None
for i in range(cfg.n_layers):
x = m.resid_lambdas[i] * x + m.x0_lambdas[i] * x0
# --- attention ---
xn = F.rms_norm(x, (D,))
# (T, H, D) - the varlen kernel's native layout, no transpose needed
q = (xn @ m.W_Q[i].mT).view(T, cfg.n_q_heads, cfg.d_qk)
k = (xn @ m.W_K[i].mT).view(T, cfg.n_kv_heads, cfg.d_qk)
v = (xn @ m.W_V[i].mT).view(T, cfg.n_kv_heads, cfg.d_vo)
# Value residual (ResFormer): value embedding mixed in via an
# input-dependent per-head gate, range (0, 3)
j = cfg.ve_index[i]
if j >= 0:
ve = F.embedding(idx, m.value_embeds[j]).view(T, cfg.n_kv_heads, cfg.d_vo)
g = 3 * torch.sigmoid(xn[..., :cfg.d_ve_gate] @ m.ve_gate[j].mT)
v = v + g.unsqueeze(-1) * ve
# Rotary embeddings (relative positional encoding)
q1, q2 = q[..., :half], q[..., half:]
k1, k2 = k[..., :half], k[..., half:]
q = torch.cat([q1 * cos + q2 * sin, q1 * (-sin) + q2 * cos], dim=-1)
k = torch.cat([k1 * cos + k2 * sin, k1 * (-sin) + k2 * cos], dim=-1)
# QK norm, then sharper attention (the 1.2 splits the scale between Q and K)
q = F.rms_norm(q, (cfg.d_qk,)) * 1.2
k = F.rms_norm(k, (cfg.d_qk,)) * 1.2
y, _ = flash_attn_varlen_fwd_lse(q, k, v, cu_seqlens, cfg.seq_len, cfg.window_sizes[i])
x = x + y.contiguous().view(T, -1) @ m.W_O[i].mT
# --- MLP (relu^2) ---
x = x + F.relu(F.rms_norm(x, (D,)) @ m.W_in[i].mT).square() @ m.W_out[i].mT
if i == backout_layer:
x_backout = x
# Subtract mid-layer residual to remove low-level features before logit projection
x = x - bf16(m.backout_lambda) * x_backout
x = F.rms_norm(x, (D,))
# lm_head + softcap
logits = (x @ m.lm_head.mT).float() # (T, d_vocab)
logits = 15.0 * torch.tanh(logits / 15.0) # smoothly cap to [-15, 15]
if targets is not None:
# No ignore_index: targets here only ever come from the training/val
# loader, which never emits pad (see forward_backward's CE note).
return F.cross_entropy(logits, targets, reduction=loss_reduction)
return logits
# --------------------------------------------------------------------------------
# § Optimizer Step
# --------------------------------------------------------------------------------
# The written-out step: one fused-kernel call per named tensor, policy at the
# call site, wrapped in the 3-phase comm flow (nanochat train_step.py):
#
# 1. Launch an async grad reduction for every sharded tensor: the full
# grad32 reduce-scatters into a fresh shard-size buffer, in the grad's
# dtype (bf16 for the two embedding tables, fp32 for everything else).
# ReduceOp.AVG across ranks composes with loss_scale=1/grad_accum_steps
# to make every reduced grad the global-batch mean.
# 2. In launch order (the comm stream completes reduces in that order):
# wait for the tensor's reduced grad, run its update kernel on the owned
# shard, then launch the async all-gather that writes the updated bf16
# live shard back into every rank's full tensor. The gather is IN PLACE
# -- our slice of the live tensor is the gather source, NCCL's
# sanctioned in-place form; even divisibility (§ Shard Assignment) means
# no padded staging buffer and no crop afterwards. Each gather overlaps
# the updates that follow it. Replicated params (ve_gate, the fp32
# scalars) ride along inline: plain all_reduce, then the identical
# full-size update on every rank.
# 3. Wait out the gathers.
#
# Waits are stream waits, not host syncs -- the whole step stays async on the
# host, and t_step still advances on-device. At world_size == 1 every
# collective short-circuits and every shard view is the whole tensor: one
# code path, degenerate comm, numerics identical to the validated single-GPU
# step.
#
# NOTE: the world>1 path has not run yet (the reference's comm code never ran
# at world>1 either) -- it awaits an 8-GPU validation pass.
ns_steps = 5 # Polar Express iterations per Muon step
# Per-slice Muon LR/WD multipliers: each bank's sqrt(max(1, fan_out/fan_in))
# aspect scale -- Muon's tall-matrix correction -- kept OUT of the shared
# matrix table so that table stays one set of numbers valid for every bank.
# At d24 only W_in is non-square, so only it gets a real multiplier (2.0).
# TODO(Chris) - I'd like to drop this eventually. If/when we drop the 2x on
# W_in we'll probably take a hit, since everything else is tuned around
# it. I don't think the trick is principled--in modded-nanogpt I
# accidentally flipped it to 2x on the mlp output and it improved loss;
# Karpathy tried that on nanochat and it didn't help. I think the model
# mostly adapts to it, so it's not worth the hassle. Get things working
# as-is first, though.
mul_unit = torch.full((cfg.n_layers, 1, 1), 1.0, dtype=torch.float32, device=device) # W_Q/W_K/W_V/W_O (square), W_out (wide -> clamped)
mul_W_in = torch.full((cfg.n_layers, 1, 1), (cfg.d_mlp / cfg.d_model) ** 0.5,
dtype=torch.float32, device=device) # 2.0 (4x expansion, tall)
mul_ve_unit = torch.full((cfg.num_ves, 1, 1), 1.0, dtype=torch.float32, device=device) # ve_gate (square)
# THE schedule position: one (1,) int64 device tensor, advanced on-device at
# the end of optimizer_step -- the host never syncs on it.
t_step = torch.zeros(1, dtype=torch.int64, device=device)
@torch.no_grad()
def optimizer_step():
"""One explicit optimizer step, written out per named tensor. Reads the
global `sched` (bind build_schedules' result to `sched` before training).
Muon MUTATES the grad it is handed (nesterov lerp) -- grad32 itself at
world=1, the reduce-scattered shard at world>1 -- so zero every grad32
afterwards either way (the loop in § Main Loop does).
no_grad is load-bearing for the fp32 scalar kernel's in-place leaf updates
(the mantissa kernels only dodge autograd's leaf check via their int
views)."""
eps = sched.adamw_eps
# ---- Phase 1: launch every async grad reduction --------------------------
# Fresh shard buffers each step (the caching allocator makes this free);
# the state tensors already carry the shard geometry, so empty_like is the
# whole allocation story.
reduced = {} # tensor -> (async work handle, shard-size reduced grad)
if world_size > 1:
for p in (m.W_Q, m.W_K, m.W_V, m.W_O, m.W_in, m.W_out):
g_shard = torch.empty_like(p.frst_mntm) # (layer shard, out, in) fp32
reduced[p] = (dist.reduce_scatter_tensor(g_shard, p.grad32, op=dist.ReduceOp.AVG, async_op=True), g_shard)
for p in (m.lm_head, m.input_embeds, m.value_embeds):
g_shard = torch.empty_like(p.exp_avg, dtype=p.grad32.dtype) # (row shard, cols) in the grad's dtype
reduced[p] = (dist.reduce_scatter_tensor(g_shard, p.grad32.view(-1, p.shape[-1]), op=dist.ReduceOp.AVG, async_op=True), g_shard)
# ---- Phase 2: wait -> owned-shard update -> gather the live shard --------
gathers = []
# Muon banks, sharded over layers
for p, mul in ((m.W_Q, mul_unit), (m.W_K, mul_unit), (m.W_V, mul_unit),
(m.W_O, mul_unit), (m.W_in, mul_W_in), (m.W_out, mul_unit)):
if world_size > 1:
work, grad = reduced[p]
work.wait()
else:
grad = p.grad32
muon_step_fused(grad, p[layer_shard_slice], p.mantissa, p.frst_mntm, p.scnd_mntm,
sched.matrix, t_step, ns_steps, p.residual_dim,
mul[layer_shard_slice], mul[layer_shard_slice])
if world_size > 1:
gathers.append(dist.all_gather_into_tensor(p, p[layer_shard_slice], async_op=True))
# Muon replicated: ve_gate is tiny, every rank updates all of it
if world_size > 1:
dist.all_reduce(m.ve_gate.grad32, op=dist.ReduceOp.AVG)
muon_step_fused(m.ve_gate.grad32, m.ve_gate, m.ve_gate.mantissa, m.ve_gate.frst_mntm, m.ve_gate.scnd_mntm, sched.matrix, t_step, ns_steps, m.ve_gate.residual_dim, mul_ve_unit, mul_ve_unit)
# AdamW, sharded over vocab rows. value_embeds' state is shaped over the
# flattened (ve_slot * vocab) row axis, so live/grad pass 2-D views
# throughout (a no-op reshape for the two already-2-D tables).
# The roles differ only in their tables (peaks/betas: build_schedules):
# lm_head runs the coolest peak (~40x below the embeddings); input_embeds
# the hottest, with the heaviest second-moment smoothing (beta2 .995);
# value_embeds rides the embedding schedule at half peak and 10x the decay.
for p, table, row_shard in ((m.lm_head, sched.lm_head, vocab_shard_slice),
(m.input_embeds, sched.input_embeds, vocab_shard_slice),
(m.value_embeds, sched.value_embeds, ve_row_shard_slice)):
rows = p.view(-1, p.shape[-1])
if world_size > 1:
work, grad = reduced[p]
work.wait()
else:
grad = p.grad32.view(-1, p.shape[-1])
adamw_step_fused(rows[row_shard], p.mantissa, grad, p.exp_avg, p.exp_avg_sq, table, t_step, eps)
if world_size > 1:
gathers.append(dist.all_gather_into_tensor(rows, rows[row_shard], async_op=True))
# AdamW replicated scalars (fp32-live, no mantissa). Three schedule
# flavors: resid -- the gentlest peak and the only decayed scalars (wd
# .05); x0 -- the hottest peak with a slow first moment (beta1 .96);
# smear -- one flat middling peak shared by all three smear/backout
# scalars, no decay. (Peaks/betas: build_schedules.)
if world_size > 1:
for p in (m.resid_lambdas, m.x0_lambdas, m.smear_gate, m.smear_lambda, m.backout_lambda):
dist.all_reduce(p.grad32, op=dist.ReduceOp.AVG)
adamw_step_fused_fp32(m.resid_lambdas, m.resid_lambdas.grad32, m.resid_lambdas.exp_avg, m.resid_lambdas.exp_avg_sq, sched.resid, t_step, eps)
adamw_step_fused_fp32(m.x0_lambdas, m.x0_lambdas.grad32, m.x0_lambdas.exp_avg, m.x0_lambdas.exp_avg_sq, sched.x0, t_step, eps)
adamw_step_fused_fp32(m.smear_gate, m.smear_gate.grad32, m.smear_gate.exp_avg, m.smear_gate.exp_avg_sq, sched.smear, t_step, eps)
adamw_step_fused_fp32(m.smear_lambda, m.smear_lambda.grad32, m.smear_lambda.exp_avg, m.smear_lambda.exp_avg_sq, sched.smear, t_step, eps)
adamw_step_fused_fp32(m.backout_lambda, m.backout_lambda.grad32, m.backout_lambda.exp_avg, m.backout_lambda.exp_avg_sq, sched.smear, t_step, eps)
# ---- Phase 3: wait out the live all-gathers ------------------------------
for work in gathers:
work.wait()
t_step.add_(1) # advance the schedule on-device
# Model + optimizer state is CAPTURED to disk at cfg.save_steps and at the
# final step (write_checkpoint, below the seam): live weights, masters via
# mantissa, both optimizers' moments, and the step counter -- world-agnostic.
# There is still deliberately no LOAD path (runs start from scratch, see the
# design decisions at the top); resume arrives with the load half when first
# needed.
##########################################################################################
# Code below comes from the 'stacks' repo
# I pulled it mainly for:
# - Pre-tokenized data, and the distributed data loader
# - Simplified (maybe?) CORE eval code
#
##########################################################################################
# --------------------------------------------------------------------------------
# § Dataset Download
# --------------------------------------------------------------------------------
NUM_TRAIN_SHARDS = 80 # full 5,568-step horizon: 70 (downloads shards 1-69,
# 6.9B raw ~= 6.1B usable after seq_len truncation --
# see the token-floor assert below the seam; 91 shards
# of 100M raw tokens are on the hub)
#DATASET_NAME = "fineweb_edu_32k_8_370"
DATASET_NAME = "climbmix_32k_8_170"
# Subdir for PT train/val .bin shards
#PT_DATA_SUBDIR = "fineweb_edu"
PT_DATA_SUBDIR = "climbmix"
HF_REPO_ID = f"ChrisMcCormick/{DATASET_NAME}"
_data_path = os.environ.get("DATA_PATH", ".")
DATASET_DIR = os.path.join(_data_path, f"data/{DATASET_NAME}")
_config_path = os.path.join(DATASET_DIR, "config.json")
train_files = os.path.join(DATASET_DIR, f"{PT_DATA_SUBDIR}/train_*.bin")
val_files = os.path.join(DATASET_DIR, f"{PT_DATA_SUBDIR}/val_*.bin")
if master_process:
from huggingface_hub import HfApi, hf_hub_download, login
hf_token = os.environ.get("HF_TOKEN")
if hf_token:
login(token=hf_token)
os.makedirs(DATASET_DIR, exist_ok=True)
api = HfApi()
train_prefix = f"{PT_DATA_SUBDIR}/train_"
to_download = []
for fname in api.list_repo_files(repo_id=HF_REPO_ID, repo_type="dataset"):
if fname.startswith(train_prefix) and int(fname[len(train_prefix):].split(".")[0]) >= NUM_TRAIN_SHARDS:
continue
if not os.path.exists(os.path.join(DATASET_DIR, fname)):
to_download.append(fname)
if to_download:
print(f"=== Downloading {len(to_download)} files from {HF_REPO_ID} ===")
for fname in to_download:
hf_hub_download(repo_id=HF_REPO_ID, filename=fname, repo_type="dataset", local_dir=DATASET_DIR)
print(" Done.")
dist.barrier()
# Load vocab config
with open(_config_path) as f:
_vocab_config = json.load(f)
VOCAB_SIZE = _vocab_config["vocab_size"]
BOS_ID = _vocab_config["bos_id"]
assert VOCAB_SIZE == cfg.d_vocab, \
f"dataset vocab ({VOCAB_SIZE}) != model d_vocab ({cfg.d_vocab}) -- wrong dataset for this hardcoded model"
# --------------------------------------------------------------------------------
# § Distributed Data Loader
# --------------------------------------------------------------------------------
# Based on the dataloader from modded-nanogpt.
# - Designed for use with flashattention_varlen_func, meaning it returns a packed token
# buffer of sequences and their lengths via cu_seqlens.
# - Hardcoded for single-epoch training.
# - Compared to `modded`, it does not support changing batch size mid-training.
def _load_data_shard(file: Path):
header = torch.from_file(str(file), False, 256, dtype=torch.int32) # header is 256 int32
assert header[0] == 20240520, "magic number mismatch in the data .bin file"
assert header[1] == 1, "unsupported version"
num_tokens = int(header[2]) # number of tokens (claimed)
with file.open("rb", buffering=0) as f:
tokens = torch.empty(num_tokens, dtype=torch.uint16, pin_memory=True) # avoid pin_memory copy by @YouJiacheng
f.seek(256 * 4)
nbytes = f.readinto(tokens.numpy()) # avoid bytes->array copy by @YouJiacheng
assert nbytes == 2 * num_tokens, "number of tokens read does not match header"
return tokens
class Shard:
def __init__(self, tokens: Tensor, world_size: int = 1):
self.tokens = tokens
self.size = tokens.numel()
self.world_size = world_size
self.i = 0
# Partial index now, full index async
self.bos_idx = (tokens[:6_000_000] == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy()
self._full_idx = None
self._loader_thread = None
self._ready = threading.Event()
self._loader_thread = threading.Thread(target=self._scan)
self._loader_thread.start()
def _scan(self):
self._full_idx = (self.tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy()
self._ready.set()
def _maybe_switch(self):
# Switch to full index as soon as async scan completes
if self.bos_idx is not self._full_idx and self._ready.is_set():
self._loader_thread.join()
self.bos_idx = self._full_idx
def next_batch(self, num_tokens_local: int, max_seq_len: int):
"""Returns (starts, ends) per rank, or None if this shard is exhausted."""
self._maybe_switch()
n = len(self.bos_idx)
starts = [[] for _ in range(self.world_size)]
ends = [[] for _ in range(self.world_size)]
idx = self.i
for r in range(self.world_size):
cur_len = 0
while cur_len <= num_tokens_local:
if idx >= n:
return None
cur = self.bos_idx[idx]
starts[r].append(cur)
end = min(self.bos_idx[idx + 1] if idx + 1 < n else self.size,
cur + max_seq_len,
cur + num_tokens_local - cur_len + 1)
ends[r].append(end)
cur_len += end - cur
idx += 1
assert cur_len == num_tokens_local + 1
self.i = idx
return starts, ends
@staticmethod
def load_async(file: Path, world_size: int = 1):
"""Returns getter function for async shard loading"""
result = {}
ready = threading.Event()
def load():
tokens = _load_data_shard(file)
result['shard'] = Shard(tokens, world_size)
ready.set()
thread = threading.Thread(target=load)
thread.start()
def get():
ready.wait()
thread.join()
return result['shard']
return get
def distributed_data_generator(filename_pattern: str, num_tokens: int, max_seq_len: int, grad_accum_steps: int = 1):
"""
Generator (i.e., yields rather than returns) of the token ids for a
micro-batch: num_tokens / grad_accum_steps / world_size tokens per yield
(32,768 for the d24 spec: total batch 2^20, grad accum 32 at world=1).
Provides both the input and target ids.
Sequences are BOS-aligned and only returned from their beginning; tokens
past max_seq_len are discarded (the next sequence starts at the next BOS).
Also used for validation batches.
Args:
filename_pattern: pattern to match the dataset .bin shard files
num_tokens: tokens per full batch (2^20 for training)
max_seq_len: 2048
grad_accum_steps: micro-batches per full batch
"""
# This GPU's rank and total GPU count.
rank = dist.get_rank() if dist.is_initialized() else 0
world_size = dist.get_world_size() if dist.is_initialized() else 1
# Confirm it all divides evenly, then calculate the per-GPU micro-batch size.
assert num_tokens % (world_size * grad_accum_steps) == 0, "Batch size must be divisible by world size"
num_tokens_local = num_tokens // grad_accum_steps // world_size
# cu_seqlens is FIXED SIZE (the compiled graph needs one shape), and ghost
# entries cost real FA3 varlen overhead, so it is sized to the DATA rather
# than a rounded guess: the densest run of climbmix docs packs 82 into one
# 32,768-token micro-batch (measured -- scan_max_docs.py; an upper bound,
# since batches can only start where the previous one ended). 96 gives
# ~17% headroom (nanochat's own estimate for these shapes also lands on
# 96), and the overflow assert below fails loudly rather than corrupt if
# the data ever changes.
max_num_docs = 192
# Get the list of shard files and wrap in an iterator.
files = [Path(file) for file in sorted(glob.glob(filename_pattern))]
if not files:
raise FileNotFoundError(f"No files found for pattern: {filename_pattern}")
file_iter = iter(files)
# Load the first shard.
tokens = _load_data_shard(next(file_iter))
shard = Shard(tokens, world_size)
remaining_files = list(file_iter)
next_shard_idx = 0
next_shard_getter = Shard.load_async(remaining_files[0], world_size) if remaining_files else None
while True:
# Get the start and end indices (within `tokens`) of the sequences to use for
# the current micro-batch.
result = shard.next_batch(num_tokens_local, max_seq_len)
# If this shard is exhausted,
if result is None:
# If there are no more shards, kill the dataloader.
if next_shard_getter is None:
return
# Load the next shard.
shard = next_shard_getter()
tokens = shard.tokens
next_shard_idx += 1
next_shard_getter = Shard.load_async(remaining_files[next_shard_idx], world_size) if next_shard_idx < len(remaining_files) else None
# Re-start the loop.
continue
# Locations of the documents in `tokens`. Only specifies the
# number of documents needed, not max.
start_idxs = torch.tensor(result[0][rank])
end_idxs = torch.tensor(result[1][rank])
# `tokens` contains the entire shard. The sequences defined by the starts and ends
# may or may not be contiguous within `tokens`, due to some sequences being
# truncated, so we slice them and then re-concatenate into a single tensor.
buf = torch.cat([tokens[i:j] for i, j in zip(start_idxs, end_idxs)])
# `buf` contains `num_tokens_local + 1` tokens to allow for the inputs vs.
# targets offset.
_inputs = buf[:-1] # All tokens minus the last
_targets = buf[1:] # Shift the tokens to the left, so that targets contains the
# next token for each input token.
# The final document includes an extra token that is the target of the last
# token in the last document. Now that we have our `_targets`, we can remove it.
end_idxs[-1] -= 1
# Calculate the start indices of the documents within `_inputs`. (flashattention
# start_idxs are relative to the `tokens` buffer, so we convert them by
# accumulating the document lengths.
# cum_lengths starts with the second document, so we'll shift
cum_lengths = (end_idxs - start_idxs).cumsum(0)
# One entry per doc plus the leading 0 must fit the fixed buffer.
assert len(cum_lengths) < max_num_docs, \
f"micro-batch packed {len(cum_lengths)} docs; cu_seqlens holds only {max_num_docs}"
# The actual cu_seqlens array always needs to contain `max_num_docs` elements so we
# the compiler can build a single graph.
# We allocate that buffer here and fill it with "empty documents", i.e., setting their start index
# to one past the end of the `_inputs` buffer.
_cum_lengths = torch.full((max_num_docs,), num_tokens_local)
# Then copy in the lengths, inserting the first document (index 0).
_cum_lengths[0] = 0
_cum_lengths[1:len(cum_lengths) + 1] = cum_lengths
# Cast to int32 / int64 on the CPU before transfer to avoid dtype conversion during .to()
_inputs = _inputs.to(dtype=torch.int32)
_targets = _targets.to(dtype=torch.int64)
_cum_lengths = _cum_lengths.to(dtype=torch.int32)
yield (
_inputs.to(device="cuda", non_blocking=True),
_targets.to(device="cuda", non_blocking=True),
_cum_lengths.to(device="cuda", non_blocking=True),
)
# Execution resumes here on the next call.
# --------------------------------------------------------------------------------
# § CORE Evaluation
# --------------------------------------------------------------------------------
# TODO - I think we can move this to a 'core_eval.py' file, I'm no longer as
# committed to the end-to-end single file approach.
"""
CORE evaluation using pre-tokenized benchmark data.
The CORE metric (from the DCLM paper, https://arxiv.org/abs/2406.11794) evaluates
a base model on in-context learning tasks using logit-based scoring (no generation).
Pre-tokenized .pt files are produced by data/core_dataset.py and loaded at eval time.
Sequences are packed into fixed-size 1D buffers with cu_seqlens marking boundaries,
enabling batched evaluation through the compiled varlen flash attention m.
"""
# -----------------------------------------------------------------------------
# Packed CORE evaluation: batch multiple examples into fixed-length 1D buffers
def pack_for_eval(sequences, buffer_size):
"""
Pack pre-tokenized sequences into fixed-size 1D buffers for batched evaluation.
Args:
sequences: list of (tokens, start_idx, end_idx, example_idx, seq_idx_within_example)
buffer_size: fixed buffer size (must be multiple of 16)
Returns:
list of dicts with keys: input_ids, cu_seqlens, metadata
"""
assert buffer_size % 16 == 0
# CORE eval sequences can be short (~50-200 tokens), so allow many more per buffer
# than training's //300 estimate. Use //8 for generous headroom (memory is negligible).
max_num_seqs = next_multiple_of_n(buffer_size // 8, n=128)
buffers = []
cur_tokens = []
cur_cu = [0]
cur_meta = []
cur_pos = 0
for tokens, start_idx, end_idx, example_idx, seq_idx in sequences:
seq_len = len(tokens)
if seq_len > buffer_size:
continue # should not happen after truncation
if cur_pos + seq_len > buffer_size:
# Finalize current buffer
_finalize_eval_buffer(buffers, cur_tokens, cur_cu, cur_meta,
buffer_size, max_num_seqs)
cur_tokens, cur_cu, cur_meta, cur_pos = [], [0], [], 0
# Track answer span in global buffer coordinates
global_start = cur_pos + start_idx
global_end = cur_pos + end_idx
cur_meta.append((example_idx, seq_idx, global_start, global_end))
cur_tokens.extend(tokens)
cur_pos += seq_len
cur_cu.append(cur_pos)
if cur_tokens:
_finalize_eval_buffer(buffers, cur_tokens, cur_cu, cur_meta,
buffer_size, max_num_seqs)
return buffers
def _finalize_eval_buffer(buffers, cur_tokens, cur_cu, cur_meta,
buffer_size, max_num_seqs):
"""Pad and finalize a packed eval buffer."""
total_packed = len(cur_tokens)
pad_count = buffer_size - total_packed
# Input tokens: packed sequences + BOS padding
input_ids = torch.full((buffer_size,), BOS_ID, dtype=torch.int32)
input_ids[:total_packed] = torch.tensor(cur_tokens, dtype=torch.int32)
# cu_seqlens: [0, end1, end2, ..., total_packed, buffer_size, buffer_size, ...]
if pad_count > 0:
cur_cu.append(buffer_size) # ghost sequence for padding region
cu_seqlens = torch.full((max_num_seqs,), buffer_size, dtype=torch.int32)
cu_seqlens[:len(cur_cu)] = torch.tensor(cur_cu, dtype=torch.int32)
buffers.append({
'input_ids': input_ids,
'cu_seqlens': cu_seqlens,
'metadata': cur_meta,
})
# TODO - The FUCK is this?? Hahaha. Typical. Screenshotting for Twitter.
@torch.no_grad()
def forward_eval_packed(input_ids, cu_seqlens):
"""
Forward a packed 1D eval buffer through the model's scoring forward.
Returns (softcapped, fp32) logits of shape (buffer_size, vocab_size).
"""
return forward(input_ids, cu_seqlens)
@torch.no_grad()
def evaluate_task_packed(task_data, buffer_size=cfg.eval_buffer_tokens):
"""Evaluate one task using pre-tokenized sequences and packed batched evaluation."""
rank = dist.get_rank() if dist.is_initialized() else 0
world_size = dist.get_world_size() if dist.is_initialized() else 1
task_type = task_data['task_type']
num_examples = task_data['num_examples']
all_sequences = task_data['sequences']
num_seqs_per_example = task_data['num_seqs_per_example']
gold_labels = task_data['gold_labels']
# Step 1: Select this rank's share of pre-tokenized sequences
rank_examples = set(range(rank, num_examples, world_size))
sequences = [
(s['tokens'], s['start_idx'], s['end_idx'], s['example_idx'], s['seq_idx'])
for s in all_sequences if s['example_idx'] in rank_examples
]
# Step 2: Pack into fixed-size buffers
packed_buffers = pack_for_eval(sequences, buffer_size)
# Step 3: Forward pass each buffer and collect per-sequence results
seq_results = {}
for buf in packed_buffers:
input_ids = buf['input_ids'].to(device)
cu_seqlens = buf['cu_seqlens'].to(device)
logits = forward_eval_packed(input_ids, cu_seqlens)
# Per-position losses: loss[j] = -log p(input_ids[j+1] | context up to j)
target_ids = torch.roll(input_ids.long(), shifts=-1)
all_losses = F.cross_entropy(logits.float(), target_ids, reduction='none')
all_predictions = logits.argmax(dim=-1)
for example_idx, seq_idx, gs, ge in buf['metadata']:
# Answer span [gs, ge): logits at [gs-1, ge-1) predict tokens at [gs, ge)
seq_results[(example_idx, seq_idx)] = {
'losses': all_losses[gs - 1 : ge - 1],
'predictions': all_predictions[gs - 1 : ge - 1],
'input_ids': input_ids[gs : ge].long(),
}
# Step 4: Evaluate per-example correctness
correct = torch.zeros(num_examples, dtype=torch.float32, device=device)
for idx in range(rank, num_examples, world_size):
if task_type == 'language_modeling':
r = seq_results[(idx, 0)]
is_correct = torch.all(r['predictions'] == r['input_ids']).item()
elif task_type in ['multiple_choice', 'schema']:
mean_losses = []
for seq_j in range(num_seqs_per_example[idx]):
r = seq_results[(idx, seq_j)]
mean_losses.append(r['losses'].mean().item())
pred_idx = mean_losses.index(min(mean_losses))
is_correct = pred_idx == gold_labels[idx]
else:
raise ValueError(f"Unsupported task type: {task_type}")
correct[idx] = float(is_correct)
if world_size > 1:
dist.barrier()
dist.all_reduce(correct, op=dist.ReduceOp.SUM)
return correct.mean().item()
@torch.no_grad()
def evaluate_chat_task_packed(task_data, buffer_size=cfg.eval_buffer_tokens):
"""Evaluate one chat categorical task using packed batched evaluation.
Unlike CORE eval (which compares losses across multiple sequences per example),
chat eval checks single-token logits at the answer position against letter choices.
Each sequence ends with the prompt (including <|assistant_start|>), and we check
what the model predicts as the next token, restricted to the valid answer letters.
"""
rank = dist.get_rank() if dist.is_initialized() else 0
world_size = dist.get_world_size() if dist.is_initialized() else 1
all_sequences = task_data['sequences']
num_examples = task_data['num_examples']
# Step 1: Select this rank's share and convert to pack_for_eval format.
# We store answer_pos as start_idx (end_idx = start_idx + 1 for tuple compat)
# and keep letter_token_ids / gold in a side table.
sequences = []
example_meta = {} # example_idx -> (letter_token_ids, gold)
for s in all_sequences:
idx = s['example_idx']
if idx % world_size != rank:
continue
answer_pos = s['answer_pos']
sequences.append((s['tokens'], answer_pos, answer_pos + 1, idx, 0))
example_meta[idx] = (s['letter_token_ids'], s['gold'])
# Step 2: Pack into fixed-size buffers (reuse CORE eval packing infrastructure)
packed_buffers = pack_for_eval(sequences, buffer_size)
# Step 3: Forward pass each buffer and score
correct = 0
total = 0
for buf in packed_buffers:
input_ids = buf['input_ids'].to(device)
cu_seqlens = buf['cu_seqlens'].to(device)
logits = forward_eval_packed(input_ids, cu_seqlens)
for example_idx, seq_idx, gs, ge in buf['metadata']:
# gs = global position of answer_pos in the buffer.
# logits[gs] predicts the token AFTER position gs — i.e. the assistant's answer.
# (This differs from CORE's logits[gs-1:ge-1] convention because here the
# answer token is NOT in the sequence — we want what the model predicts next.)
answer_logits = logits[gs] # (vocab_size,)
letter_ids, gold = example_meta[example_idx]
focus_logits = answer_logits[letter_ids] # (num_choices,)
pred = focus_logits.argmax().item()
correct += int(pred == gold)
total += 1
# Step 4: Aggregate across ranks
if world_size > 1:
correct_t = torch.tensor([correct], dtype=torch.long, device=device)
total_t = torch.tensor([total], dtype=torch.long, device=device)
dist.all_reduce(correct_t, op=dist.ReduceOp.SUM)
dist.all_reduce(total_t, op=dist.ReduceOp.SUM)
correct = correct_t.item()
total = total_t.item()
return correct / total if total > 0 else 0.0
def evaluate_chat_categorical():
"""
Evaluate a chat model on categorical benchmarks (MMLU, ARC-Easy, ARC-Challenge)
using pre-tokenized data from chat_eval_dataset.py.
Returns dict with results, centered_results, and chatcore_metric.
"""
chat_eval_dir = os.path.join(DATASET_DIR, "chat_eval")
config_path = os.path.join(chat_eval_dir, "config.json")
assert os.path.exists(config_path), f"Chat eval config not found: {config_path}"
with open(config_path, 'r', encoding='utf-8') as f:
config = json.load(f)
# Evaluate each task
results = {}
centered_results = {}
for task_info in config['tasks']:
torch.cuda.synchronize()
start_time = time.time()
label = task_info['label']
pt_path = os.path.join(chat_eval_dir, task_info['file'])
assert os.path.exists(pt_path), f"Chat eval data not found: {pt_path}"
task_data = torch.load(pt_path, weights_only=False)
print0(f"Chat eval: {label} ({task_data['num_examples']} examples)... ", console=True)
accuracy = evaluate_chat_task_packed(task_data)
torch.cuda.synchronize()
results[label] = accuracy
random_baseline = task_data['random_baseline']
centered_result = (accuracy - random_baseline) / (1.0 - random_baseline)
centered_results[label] = centered_result
elapsed = time.time() - start_time
print0(f"accuracy: {accuracy:.4f} | centered: {centered_result:.4f} | time: {elapsed:.2f}s", console=True)
chatcore_metric = sum(centered_results.values()) / len(centered_results)
out = {
"results": results,
"centered_results": centered_results,
"chatcore_metric": chatcore_metric,
}
return out
def evaluate_core():
"""
Evaluate a base model on the CORE benchmark using pre-tokenized data.
Returns dict with results, centered_results, and core_metric.
"""
core_eval_dir = os.path.join(DATASET_DIR, "core_eval")
config_path = os.path.join(core_eval_dir, "config.json")
with open(config_path, 'r', encoding='utf-8') as f:
config = json.load(f)
# Evaluate each task
results = {}
centered_results = {}
for task_info in config['tasks']:
torch.cuda.synchronize()
start_time = time.time()
label = task_info['label']
task_data = torch.load(os.path.join(core_eval_dir, task_info['file']),
weights_only=False)
print0(f"Evaluating: {label} ({task_data['task_type']}, "
f"{task_data['num_examples']} examples)... ", console=True)
accuracy = evaluate_task_packed(task_data)
torch.cuda.synchronize()
results[label] = accuracy
random_baseline = task_data['random_baseline']
centered_result = (accuracy - 0.01 * random_baseline) / (1.0 - 0.01 * random_baseline)
centered_results[label] = centered_result
elapsed = time.time() - start_time
print0(f"accuracy: {accuracy:.4f} | centered: {centered_result:.4f} | time: {elapsed:.2f}s", console=True)
core_metric = sum(centered_results.values()) / len(centered_results)
out = {
"results": results,
"centered_results": centered_results,
"core_metric": core_metric
}
return out
# --------------------------------------------------------------------------------
# § Main Loop
# --------------------------------------------------------------------------------
# Modeled on nanochat base_train (branch fwd-bwd) -- the flat trainer over the
# same forward_backward / optimizer_step API. No warmup-and-reset phase (that
# trick needs the state_dict save/restore this file deliberately lacks):
# compilation happens during the first real steps, and the time totals simply
# exclude the first 10 steps (the nanochat convention).
# begin logging
logfile = None
if master_process:
run_id = cfg.run_id
os.makedirs("logs", exist_ok=True)
logfile = f"logs/{run_id}.txt"
print(logfile)
def print0(s="", console=False):
if master_process:
with open(logfile, "a") as f:
if console:
print(s)
print(s, file=f)
print0(code)
print0("="*100)
print0(f"Running Python {sys.version}")
print0(f"Running PyTorch {torch.version.__version__} compiled for CUDA {torch.version.cuda}")
# -----------------------------------------------------------------------------
# Model stats, for MFU and the wandb config -- CONSTANTS; scaling.py recomputes
# them from the d24 shapes (params by group, 6 FLOPs per matmul-weight param
# plus the windowed attention term).
num_params = 1_384_122_122 # every trained weight (Model.weight_names)
num_flops_per_token = 4_860_160_128 # 6 * 729,810,624 matmul params + attention
gpu_device_name = torch.cuda.get_device_name(0)
gpu_peak_flops = next((v for k, v in PEAK_FLOPS.items() if k in gpu_device_name.upper()),
float("inf"))
print0(f"Model parameters: {num_params:,} | FLOPs/token: {num_flops_per_token:e}", console=True)
print0(f"GPU: {gpu_device_name} | Peak FLOPS (BF16): {gpu_peak_flops:.2e}", console=True)
print0(f"Total batch size: {cfg.total_batch_size:,} tokens = {cfg.micro_batch_tokens:,} tokens/micro "
f"x {world_size} ranks x {grad_accum_steps} grad accum", console=True)
# -----------------------------------------------------------------------------
# Schedules: every LR/beta/WD coefficient for the whole run, materialized up
# front. The two batch/horizon corrections are hardcoded (derivations in
# scaling.py):
# batch_lr_scale = sqrt(2^20 / 2^19) = 1.4142... -- eta ∝ sqrt(B/B_ref),
# B_ref = 2^19 where the d12 LRs were tuned; build_schedules applies it to
# the per-role peaks itself (do NOT also fold it into the LRs).
# weight_decay = 0.28 * sqrt(2) * (d12/d24 scaling params) = 0.059738 -- the
# T_epoch framework; matches nanochat's d24 printout exactly.
sched = build_schedules(cfg.num_iterations, batch_lr_scale=1.4142135623730951,
weight_decay=0.059738)
# -----------------------------------------------------------------------------
# Compile the training step. REQUIRED, not an optimization: the CE block in
# forward_backward is written for inductor's fusion -- run eager it
# materializes full (T, d_vocab) fp32 temporaries. fullgraph so any graph
# break errors loudly instead of silently fragmenting fusion (the FA3 raw ops
# have fake impls, so a full trace is achievable).
fb = torch.compile(forward_backward, dynamic=False, fullgraph=True)
# The eval forward is compiled too -- eager it materializes the full
# (T, d_vocab) fp32 logits chain (~13 GB of temporaries per val micro-batch).
# Rebinding the name routes every consumer (the val-loss section and
# forward_eval_packed) through it; it specializes once per shape/targets
# combination: the val path at the training micro-batch shape, the CORE
# logits path at the eval buffer shape.
forward = torch.compile(forward, dynamic=False, fullgraph=True)
# token_bytes: per-token-id byte lengths (0 for special tokens), for the
# vocab-size-independent bits-per-byte validation metric.
with open(os.path.join(DATASET_DIR, "tokenizer/token_bytes.pt"), "rb") as f:
token_bytes = torch.load(f, map_location=device)
# Enough data for the horizon? The loader is single-epoch and TRUNCATES long
# documents at seq_len, discarding the tails: measured ~11-12% of climbmix's
# raw tokens (doc-length scan, 2026-07-31 session NOTES). 0.85 is that
# discard with margin -- a raw-token floor alone would pass configs that run
# dry ~11% before the horizon.
_shard_tokens = sum((os.path.getsize(f) - 256 * 4) // 2 for f in glob.glob(train_files))
assert _shard_tokens * 0.85 >= (cfg.num_iterations + 1) * cfg.total_batch_size, \
f"train shards hold {_shard_tokens:,} raw tokens (~{int(_shard_tokens * 0.85):,} usable " \
f"after seq_len truncation) < {(cfg.num_iterations + 1) * cfg.total_batch_size:,} needed " \
f"-- raise NUM_TRAIN_SHARDS"
# --- wandb logging init ---
use_dummy_wandb = cfg.wandb_run == "dummy" or not master_process
wandb_run = DummyWandb() if use_dummy_wandb else wandb.init(
project=cfg.wandb_project, name=cfg.wandb_run,
config={
"num_params": num_params,
"num_flops_per_token": num_flops_per_token,
"n_layers": cfg.n_layers, "n_q_heads": cfg.n_q_heads, "d_model": cfg.d_model,
"train_steps": cfg.num_iterations,
"total_batch_size": cfg.total_batch_size,
"micro_batch_tokens": cfg.micro_batch_tokens,
"val_loss_every": cfg.val_loss_every,
"world_size": world_size,
"grad_accum_steps": grad_accum_steps,
"batch_lr_scale": sched.batch_lr_scale,
"weight_decay": sched.weight_decay,
},
)
if not use_dummy_wandb:
wandb.define_metric("step")
wandb.define_metric("*", step_metric="step")
# -----------------------------------------------------------------------------
# Checkpoint capture (write only -- there is deliberately no load/resume path
# yet). Two files per capture point in logs/{run_id}/:
# model_stepNNNNNN.pt -- {step, code, weights: {name: tensor}} -- the bf16
# live weights + fp32 scalars, the payload the final save has always held.
# optim_stepNNNNNN.pt -- {step, t_step, state: {"name.attr": tensor}} over
# the five optimizer-state attrs; together with the live weights this is
# the full fp32 masters and both optimizers' moments.
# World-agnostic: sharded state all-gathers to full size before writing, so a
# capture from an 8-GPU run loads at any world size (at world=1 the gathers
# short-circuit and this is a plain copy-out). Every rank participates in the
# gathers; only master materializes CPU copies and writes -- tensors are saved
# on CPU so the files open anywhere.
state_attrs = ("mantissa", "frst_mntm", "scnd_mntm", "exp_avg", "exp_avg_sq")
# The sharded weights -- their state gathers over dim 0; everything else is
# replicated, already full-size on every rank. Mirrors § Shard Assignment.
# A set, not a tuple: tuple membership falls through identity to elementwise
# tensor ==, while set membership stays on the identity hash.
sharded_weights = {m.W_Q, m.W_K, m.W_V, m.W_O, m.W_in, m.W_out,
m.lm_head, m.input_embeds, m.value_embeds}
def gather_full(t):
"""All-gather a shard-size state tensor to full size over dim 0. uint16
(mantissa) rides as a bf16 bitcast: NCCL has no 16-bit int type, and a
gather only moves bytes."""
if world_size == 1:
return t
comm = t.view(torch.bfloat16) if t.dtype == torch.uint16 else t
full = torch.empty(t.shape[0] * world_size, *t.shape[1:], dtype=comm.dtype, device=device)
dist.all_gather_into_tensor(full, comm)
return full.view(torch.uint16) if t.dtype == torch.uint16 else full
def write_checkpoint(step):
state = {}
for n in m.weight_names:
p = getattr(m, n)
for attr in state_attrs:
if hasattr(p, attr):
full = gather_full(getattr(p, attr)) if p in sharded_weights else getattr(p, attr)
if master_process:
state[f"{n}.{attr}"] = full.cpu()
if not master_process:
return
os.makedirs(f"logs/{run_id}", exist_ok=True)
torch.save(dict(step=step, code=code,
weights={n: getattr(m, n).cpu() for n in m.weight_names}),
f"logs/{run_id}/model_step{step:06d}.pt")
torch.save(dict(step=step, t_step=int(t_step.item()), state=state),
f"logs/{run_id}/optim_step{step:06d}.pt")
# -----------------------------------------------------------------------------
# Training and validation
train_steps = cfg.num_iterations
train_loader = distributed_data_generator(train_files, cfg.total_batch_size, cfg.seq_len, grad_accum_steps)
inputs, targets, cu_seqlens = next(train_loader) # kick off the first batch
# Each val pass draws val_tokens through micro-batches shaped exactly like
# training's (so the rotary-cache bound holds), scored with the eager forward.
micro_world_tokens = cfg.total_batch_size // grad_accum_steps # tokens per micro-batch across ranks
assert cfg.val_tokens % micro_world_tokens == 0
val_steps = cfg.val_tokens // micro_world_tokens
val_bpb = None
min_val_bpb = float("inf")
smooth_train_loss = 0.0
total_training_time = 0.0 # seconds; excludes the first 10 steps (compile lives there)
for step in range(train_steps + 1):
last_step = (step == train_steps)
# --------------- VALIDATION SECTION -----------------
if last_step or (cfg.val_loss_every > 0 and step % cfg.val_loss_every == 0):
torch.cuda.synchronize()
val_t0 = time.perf_counter()
val_loader = distributed_data_generator(val_files, cfg.total_batch_size, cfg.seq_len, grad_accum_steps)
total_nats = torch.tensor(0.0, dtype=torch.float32, device=device)
total_bytes = torch.tensor(0, dtype=torch.int64, device=device)
for _ in range(val_steps):
v_inputs, v_targets, v_cu_seqlens = next(val_loader)
loss_flat = forward(v_inputs, v_cu_seqlens, v_targets, loss_reduction='none')
num_bytes_flat = token_bytes[v_targets]
total_nats += (loss_flat * (num_bytes_flat > 0)).sum()
total_bytes += num_bytes_flat.sum()
del val_loader
if world_size > 1:
dist.all_reduce(total_nats, op=dist.ReduceOp.SUM)
dist.all_reduce(total_bytes, op=dist.ReduceOp.SUM)
val_bpb = total_nats.item() / (math.log(2) * total_bytes.item())
min_val_bpb = min(min_val_bpb, val_bpb)
val_elapsed = time.perf_counter() - val_t0
print0(f"step:{step}/{train_steps} val_bpb:{val_bpb:.6f} val_time:{val_elapsed:.2f}s", console=True)
wandb_run.log({"step": step, "val/bpb": val_bpb, "val/eval_seconds": val_elapsed,
"total_training_time": total_training_time})
# --------------- CHECKPOINT CAPTURE -----------------
# State on entering step `step` = after `step` completed updates. Every
# rank enters (the gathers are collectives); only master writes.
if cfg.save_checkpoint and (last_step or step in cfg.save_steps):
ckpt_t0 = time.perf_counter()
write_checkpoint(step)
print0(f"checkpoint captured at step {step} ({time.perf_counter() - ckpt_t0:.1f}s)", console=True)
if last_step:
# --------------- CORE EVALUATION -----------------
if os.path.exists(os.path.join(DATASET_DIR, "core_eval/config.json")):
core_eval_t0 = time.perf_counter()
core_out = evaluate_core()
core_eval_elapsed = time.perf_counter() - core_eval_t0
print0(f"CORE metric: {core_out['core_metric']:.4f} | total CORE eval time: {core_eval_elapsed:.2f}s", console=True)
for label, acc in core_out['results'].items():
print0(f" {label}: accuracy={acc:.4f} centered={core_out['centered_results'][label]:.4f}", console=True)
wandb_run.log({
"step": step,
"core_metric": core_out["core_metric"],
**{f"core/{label}/accuracy": acc for label, acc in core_out["results"].items()},
**{f"core/{label}/centered": c for label, c in core_out["centered_results"].items()},
"timing/core_eval_seconds": core_eval_elapsed,
})
else:
print0("No core_eval/ in the dataset dir; skipping the CORE metric.", console=True)
break
# --------------- TRAINING SECTION -----------------
torch.cuda.synchronize()
step_t0 = time.perf_counter()
for micro in range(grad_accum_steps):
# loss_scale replaces the loss/grad_accum division of an autograd loop
loss = fb(inputs, targets, cu_seqlens, loss_scale=grad_scale)
inputs, targets, cu_seqlens = next(train_loader) # prefetch while the GPU is busy
optimizer_step() # schedules pre-computed; advances t_step on-device
# Zero every grad buffer: gradients accumulate across the next step's
# micro-batches, and at world=1 Muon's nesterov lerp just MUTATED grad32
# (at world>1 it mutates the reduce-scattered shard instead) -- this is
# correctness, not hygiene. (`for p in m` = every trained weight, in
# Model.weight_names order.)
for p in m:
p.grad32.zero_()
train_loss = loss.item() # the step's one host sync point
torch.cuda.synchronize()
dt = time.perf_counter() - step_t0
# logging (CPU only). EMA the loss for readability; time totals exclude the
# first 10 steps, where compilation dominates.
ema_beta = 0.9
smooth_train_loss = ema_beta * smooth_train_loss + (1 - ema_beta) * train_loss
debiased_smooth_loss = smooth_train_loss / (1 - ema_beta ** (step + 1))
if step > 10:
total_training_time += dt
tok_per_sec = int(cfg.total_batch_size / dt)
mfu = 100 * num_flops_per_token * cfg.total_batch_size / dt / (gpu_peak_flops * world_size)
steps_timed = step - 10
if steps_timed > 0:
eta_seconds = (train_steps - step - 1) * (total_training_time / steps_timed)
eta_str = f" | eta: {eta_seconds/60:.1f}m"
else:
eta_str = ""
pct_done = 100 * step / train_steps
print0(f"step {step:05d}/{train_steps:05d} ({pct_done:.2f}%) | loss: {debiased_smooth_loss:.6f} | lrm: {sched.lrm_table[step]:.2f} | dt: {dt*1000:.2f}ms | tok/sec: {tok_per_sec:,} | bf16_mfu: {mfu:.2f} | total time: {total_training_time/60:.2f}m{eta_str}", console=True)
wandb_run.log({
"step": step,
"train/loss": debiased_smooth_loss,
"train/lrm": float(sched.lrm_table[step]),
"train/dt": dt,
"train/tok_per_sec": tok_per_sec,
"train/mfu": mfu,
"total_training_time": total_training_time,
})
# GC management: the collector's cycle scans cost ~500ms at random steps,
# so collect the setup garbage once, then freeze survivors and disable.
if step == 0:
gc.collect()
gc.freeze()
gc.disable()
elif step % 5000 == 0:
gc.collect()
print0(f"peak memory allocated: {torch.cuda.max_memory_allocated() // 1024 // 1024} MiB "
f"reserved: {torch.cuda.max_memory_reserved() // 1024 // 1024} MiB", console=True)
print0(f"total training time: {total_training_time/60:.2f}m", console=True)
if val_bpb is not None:
print0(f"minimum validation bpb: {min_val_bpb:.6f}", console=True)
wandb_run.finish()
dist.destroy_process_group()
|