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
|
"""BP-train a small cascade-form standard transformer (L distinct blocks), saving ckpts
every --save_every for the A0.2 on-trajectory gradient gate (cascade_probe.py --ckpt).
Plain LLM training — this is also the BP twin for the C-tier money runs."""
import argparse, math, pickle, time, json
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from pathlib import Path
ap = argparse.ArgumentParser()
ap.add_argument('--tag', default='casc_bp6')
ap.add_argument('--L', type=int, default=6); ap.add_argument('--C', type=int, default=256)
ap.add_argument('--H', type=int, default=8); ap.add_argument('--T', type=int, default=256)
ap.add_argument('--B', type=int, default=24); ap.add_argument('--steps', type=int, default=4000)
ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', type=int, default=200)
ap.add_argument('--seed', type=int, default=0)
ap.add_argument('--save_every', type=int, default=500); ap.add_argument('--log', type=int, default=200)
ap.add_argument('--wandb', default=''); ap.add_argument('--wandb_run', default='')
ap.add_argument('--opt', choices=['adamw', 'muon'], default='adamw')
ap.add_argument('--muon_lr', type=float, default=0.02)
ap.add_argument('--tok_init', type=float, default=0.0) # >0: init tok/pos std (GPT-standard 0.02)
ap.add_argument('--cosine', action='store_true') # warmup then cosine decay to lr_min_ratio*lr over --steps (long runs)
ap.add_argument('--lr_min_ratio', type=float, default=0.1)
args = ap.parse_args()
torch.manual_seed(args.seed)
dev = 'cuda' if torch.cuda.is_available() else 'cpu'
DD = Path('/home/yurenh2/ept/ep_run/data/tinystories_bpe')
vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size']
def get_batch(split):
data = np.memmap(DD / ('train.bin' if split == 'train' else 'val.bin'), dtype=np.uint16, mode='r')
ix = torch.randint(len(data) - args.T - 1, (args.B,))
x = torch.stack([torch.from_numpy(data[i:i + args.T].astype(np.int64)) for i in ix])
y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + args.T].astype(np.int64)) for i in ix])
return x.to(dev), y.to(dev)
class Block(nn.Module):
def __init__(self, C, H):
super().__init__()
self.ln1, self.ln2 = nn.LayerNorm(C), nn.LayerNorm(C)
self.attn = nn.MultiheadAttention(C, H, batch_first=True)
self.ff = nn.Sequential(nn.Linear(C, 4 * C), nn.GELU(), nn.Linear(4 * C, C))
def forward(self, z, mask):
h = self.ln1(z); a, _ = self.attn(h, h, h, attn_mask=mask, need_weights=False)
z = z + a; return z + self.ff(self.ln2(z))
tok = nn.Embedding(vocab, args.C).to(dev)
pos = nn.Embedding(args.T, args.C).to(dev)
if args.tok_init > 0:
with torch.no_grad():
tok.weight.normal_(0, args.tok_init); pos.weight.normal_(0, args.tok_init)
blocks = nn.ModuleList([Block(args.C, args.H) for _ in range(args.L)]).to(dev)
mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1)
params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters())
if args.opt == 'muon':
from muon import build_hybrid
opt, sched = build_hybrid(blocks, params, args.lr, args.muon_lr, args.warmup)
else:
opt = torch.optim.AdamW(params, lr=args.lr, weight_decay=1e-4)
if args.cosine:
def _lrlam(s):
if s < args.warmup: return (s + 1) / max(args.warmup, 1)
p = min(1.0, (s - args.warmup) / max(1, args.steps - args.warmup))
return args.lr_min_ratio + 0.5 * (1 - args.lr_min_ratio) * (1 + math.cos(math.pi * p))
sched = torch.optim.lr_scheduler.LambdaLR(opt, _lrlam)
else:
sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / max(args.warmup, 1)))
def fwd(x):
z = tok(x) + pos(torch.arange(args.T, device=dev))[None]
for b in blocks: z = b(z, mask)
return z @ tok.weight.t()
@torch.no_grad()
def evaluate(nb=6):
tot = 0.0
for _ in range(nb):
x, y = get_batch('val')
tot += F.cross_entropy(fwd(x).reshape(-1, vocab), y.reshape(-1)).item()
return tot / nb
wb = None
if args.wandb:
try:
import wandb as _w
wb = _w.init(project=args.wandb, name=args.wandb_run or args.tag, id=args.wandb_run or args.tag,
resume='allow', config=vars(args))
except Exception as e:
print(f'[wandb] disabled ({e})', flush=True)
n = sum(p.numel() for p in params)
print(f'[{args.tag}] cascade-BP L{args.L} C{args.C} H{args.H} T{args.T} | {n/1e6:.2f}M params | {dev}', flush=True)
best, t0 = 1e9, time.time()
outdir = Path('runs'); outdir.mkdir(exist_ok=True)
for step in range(args.steps + 1):
x, y = get_batch('train')
loss = F.cross_entropy(fwd(x).reshape(-1, vocab), y.reshape(-1))
opt.zero_grad(set_to_none=True); loss.backward()
torch.nn.utils.clip_grad_norm_(params, 1.0)
opt.step(); sched.step()
if step % args.log == 0:
val = evaluate(); best = min(best, val)
print(f'step {step:5d}/{args.steps} | train {loss.item():.4f} val {val:.4f} (best {best:.4f}) '
f'| {step/max(time.time()-t0,1e-9):.2f} it/s', flush=True)
if wb is not None:
try: wb.log({'train_ce': loss.item(), 'val_ce': val, 'best': best}, step=step)
except Exception: pass
if step % args.save_every == 0:
torch.save({'tok': tok.state_dict(), 'pos': pos.state_dict(), 'blocks': blocks.state_dict(),
'step': step, 'val': best, 'config': vars(args)}, outdir / f'{args.tag}_s{step}.pt')
print(f'[{args.tag}] DONE best val CE {best:.4f} (random ln({vocab})={math.log(vocab):.3f})', flush=True)
if wb is not None:
try: wb.summary['best_val_ce'] = best; wb.finish()
except Exception: pass
|