File size: 12,311 Bytes
224edbf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6f7a03c
 
224edbf
 
 
 
 
 
 
 
6f7a03c
 
 
224edbf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6f7a03c
 
 
 
 
 
 
 
 
 
224edbf
 
 
 
 
6f7a03c
224edbf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6f7a03c
224edbf
 
6f7a03c
224edbf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# train_peer_v21.py - pretrain the v21 shared-pool ternary+PEER controller.
# Same data/codon/Muon/WSD machinery as train_peer.py, with the v21 fixes:
#   P1 3 shared PEER layers   P2 embedding_bag   P4 SparseAdam on the pool
#   P5 checkpoint ONLY the PEER blocks   P6 grad-clip excludes the sparse pool
# Logs a throughput series to <out>/throughput.jsonl for the A/B charts.
import argparse, json, math, os, time, threading, queue
from pathlib import Path
import numpy as np, torch, torch.nn.functional as F
import torch.utils.checkpoint as cp
import sys; sys.path.insert(0, '/root/dna')
from model_dna_peer_v21 import DnaPeerV21, peer_counts, param_groups, compile_

ap = argparse.ArgumentParser()
ap.add_argument('--data-dir', default='/root/dna/data')
ap.add_argument('--out', default='/root/dna/ckpt-peer21'); ap.add_argument('--resume', default='/root/dna/ckpt-peer21/resume.pt')
ap.add_argument('--batch', type=int, default=16); ap.add_argument('--seq', type=int, default=512)
ap.add_argument('--steps', type=int, default=2000000); ap.add_argument('--target-tokens', type=float, default=3e9)
ap.add_argument('--max-hours', type=float, default=0.0, help='stop after N hours (0=off)')
ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--muonlr', type=float, default=2e-2)
ap.add_argument('--poollr', type=float, default=0.02)  # stateless SGD wants a larger lr
ap.add_argument('--nk', type=int, default=704); ap.add_argument('--topk', type=int, default=8)
ap.add_argument('--pheads', type=int, default=1)
ap.add_argument('--peer-layers', default='4,12')  # 2 layers MEASURED 38.2k vs 3 layers 35.7k
ap.add_argument('--rowstep', type=float, default=1.0); ap.add_argument('--warmup', type=int, default=2000)
ap.add_argument('--ckpt-every', type=int, default=2000); ap.add_argument('--hf-every', type=int, default=20000)
ap.add_argument('--hf-repo', default='jaivial/dna-diskchat-2b-peer-v21'); ap.add_argument('--log-every', type=int, default=100)
ap.add_argument('--mtp-frac', type=float, default=0.25)
a = ap.parse_args(); dev = 'cuda'; Path(a.out).mkdir(parents=True, exist_ok=True); torch.manual_seed(1)
torch._dynamo.config.cache_size_limit = 64
torch.backends.cuda.matmul.allow_tf32 = True; torch.backends.cudnn.allow_tf32 = True; torch.backends.cudnn.benchmark = True

m = DnaPeerV21(nk=a.nk, topk=a.topk, pheads=a.pheads,
               peer_layers=tuple(int(x) for x in a.peer_layers.split(',')),
               chunk=int(os.environ.get('CHUNK_C', '16'))).to(dev)
muon_p, sparse_p, dense_p = param_groups(m)
# P5: compile + checkpoint only where it pays
compile_(m)   # compile all but the sparse gathers (inductor cannot lower them)
PEER_SET = set(m.peer_layers)
ctrl, active = peer_counts(m)
CFG = {'ctrl_params': ctrl, 'active_params_per_token': active, 'experts': m.pool.E,
       'peer_layers': list(m.peer_layers), 'h_times_k': a.pheads * a.topk, 'nk': a.nk,
       'pool_params': sum(p.numel() for p in sparse_p), 'pool_dtype': str(m.pool.up.weight.dtype),
       'codon_table_bytes': m.nrows * m.ncod, 'batch': a.batch, 'seq': a.seq, 'variant': 'v21-shared-pool'}
print('PEER21_CONFIG', json.dumps(CFG), flush=True)

def ns5(Gm, steps=int(os.environ.get('NS_STEPS', '3'))):
    x = 3.4445; y = -4.7750; z = 2.0315; X = Gm.bfloat16(); tp = Gm.size(0) > Gm.size(1)
    if tp: X = X.t()
    X = X / (X.norm() + 1e-7)
    for _ in range(steps):
        A = X @ X.t(); B = y * A + z * (A @ A); X = x * X + B @ X
    if tp: X = X.t()
    return X.to(Gm.dtype)
class Muon(torch.optim.Optimizer):
    def __init__(self, params, lr=2e-2, momentum=0.95): super().__init__(list(params), dict(lr=lr, momentum=momentum))
    @torch.no_grad()
    def step(self):
        for grp in self.param_groups:
            for p in grp['params']:
                if p.grad is None: continue
                st = self.state[p]
                if 'm' not in st: st['m'] = torch.zeros_like(p.grad)
                buf = st['m']; buf.mul_(grp['momentum']).add_(p.grad); u = p.grad.add(buf, alpha=grp['momentum'])
                u = ns5(u); u = u.mul(max(1.0, p.size(0) / p.size(1)) ** 0.5); p.add_(u, alpha=-grp['lr'])

opt = torch.optim.AdamW(dense_p, lr=a.lr, betas=(.9, .95), weight_decay=.1, fused=True)
muon = Muon(muon_p, lr=a.muonlr)
# P4: STATELESS sparse SGD on the expert pool (v8 lever T5). SparseAdam was tried
# first and OOMed: it materializes dense exp_avg/exp_avg_sq (4.06 GB) plus a
# sparse_mask temporary, defeating the whole point. This keeps ZERO optimizer state.
@torch.no_grad()
def sparse_sgd(lr):
    for p in sparse_p:
        g = p.grad
        if g is None: continue
        if g.is_sparse:
            g = g.coalesce()
            p.data.index_add_(0, g.indices()[0], g.values(), alpha=-lr)
        else:
            p.data.add_(g, alpha=-lr)
        p.grad = None
print('MUON', sum(p.numel() for p in muon_p), 'ADAM', sum(p.numel() for p in dense_p),
      'SPARSE_POOL', sum(p.numel() for p in sparse_p), flush=True)

cur = None; curi = -1
def shard_list(): return sorted([p for p in Path(a.data_dir).glob('shard_*.u16') if not p.with_suffix('.tmp').exists()])
def batch():
    global cur, curi
    sh = shard_list()
    if not sh: raise SystemExit('no shards yet - start stream_data.py first')
    i = np.random.randint(0, len(sh))
    if i != curi: cur = np.memmap(sh[i], np.uint16, 'r'); curi = i
    n = len(cur) - a.seq - 1; ix = np.random.randint(0, n, size=a.batch)
    return torch.from_numpy(np.stack([np.asarray(cur[j:j+a.seq+1], np.int64) for j in ix])).to(dev)
def cosf(s):
    if s < a.warmup: return s / a.warmup
    ds = int(a.steps * 0.8)
    if s < ds: return 1.0
    p = (s - ds) / max(1, a.steps - ds); return 1.0 - 0.9 * p
def feats(ids):
    B, T = ids.shape; x = m.embed(ids)
    for i, blk in enumerate(m.blocks):
        # P5: recompute only the PEER blocks; the other 13 run without recompute
        x = cp.checkpoint(blk, x, m.chunk, use_reentrant=False) if i in PEER_SET else blk(x, m.chunk)
    mem, bal, rows = m.route(x.reshape(B * T, m.d)); return m.norm(x + mem.reshape(B, T, m.d)), bal, rows
def save(step):
    tmp = a.resume + '.tmp'
    torch.save({'model': m.state_dict(), 'opt': opt.state_dict(), 'muon': muon.state_dict(),
                'step': step, 'config': m.config()}, tmp)
    os.replace(tmp, a.resume)
def hf_push(step):
    try:
        from huggingface_hub import HfApi
        api = HfApi(token=os.environ.get('HUGGING_FACE')); api.create_repo(a.hf_repo, exist_ok=True)
        p = f'{a.out}/ckpt_{step}.pt'
        torch.save({'model': m.state_dict(), 'config': m.config(), 'step': step}, p)
        api.upload_file(path_or_fileobj=p, path_in_repo=f'checkpoints/ckpt_{step}.pt', repo_id=a.hf_repo)
        os.remove(p)
        tp = Path(a.out) / 'throughput.jsonl'
        if tp.exists(): api.upload_file(path_or_fileobj=str(tp), path_in_repo='logs/throughput.jsonl', repo_id=a.hf_repo)
        print(f'HF_PUSH step {step}', flush=True)
    except Exception as e: print('HF_PUSH_FAIL', str(e)[:140], flush=True)

def chunked_ce(feat, W, tgt, nchunk=8):
    """Tiled cross-entropy (v12 lever A). Never materializes the full [N,V] logit
    tensor; mathematically identical to F.cross_entropy(F.linear(feat,W),tgt).
    MEASURED: peak VRAM 11.32 -> 10.56 GB at B20, which is what pays for B28."""
    N = feat.size(0); tot = feat.new_zeros(())
    for f, t in zip(feat.chunk(nchunk, 0), tgt.chunk(nchunk, 0)):
        tot = tot + F.cross_entropy(F.linear(f, W), t, reduction='sum')
    return tot / N


@torch.no_grad()
def valid_loss():
    m.eval(); ids = batch()
    with torch.autocast('cuda', dtype=torch.bfloat16):
        feat, _, _ = feats(ids)
        ce = chunked_ce(feat[:, :-1].reshape(-1, m.d), m.embed.weight, ids[:, 1:].reshape(-1))
    m.train(); return ce.item()

start = 0
if os.path.exists(a.resume):
    ck = torch.load(a.resume, map_location=dev, weights_only=False)
    m.load_state_dict(ck['model']); opt.load_state_dict(ck['opt']); muon.load_state_dict(ck['muon'])
    start = int(ck['step']); print('RESUME', start, flush=True)

_bq = queue.Queue(maxsize=4)
def _worker():
    while True: _bq.put(batch())
for _ in range(2): threading.Thread(target=_worker, daemon=True).start()
def nextbatch(): return _bq.get()

# ---- throughput tracking (the A/B deliverable) ----
TP = Path(a.out) / 'throughput.jsonl'
tp_f = open(TP, 'a')
ts_hist = []
def log_tp(rec):
    tp_f.write(json.dumps(rec) + '\n'); tp_f.flush()

m.train(); t0 = time.time(); wall0 = time.time(); seen = 0; step = start
print('LOOP start', start, 'steps', a.steps, flush=True)
for step in range(start + 1, a.steps + 1):
    f = cosf(step)
    for g in opt.param_groups: g['lr'] = a.lr * f
    for g in muon.param_groups: g['lr'] = a.muonlr * f
    ids = nextbatch()
    opt.zero_grad(set_to_none=True)
    for p in sparse_p: p.grad = None
    for p in muon_p:
        p.grad = None
    with torch.autocast('cuda', dtype=torch.bfloat16):
        feat, bal, rows = feats(ids)
        fl = feat[:, :-1].reshape(-1, m.d); tg = ids[:, 1:].reshape(-1)
        l1 = chunked_ce(fl, m.embed.weight, tg)
        f2 = feat[:, :-2].reshape(-1, m.d); t2 = ids[:, 2:].reshape(-1)
        sel = torch.randint(0, f2.size(0), (int(a.mtp_frac * f2.size(0)),), device=dev)
        l2 = chunked_ce(f2[sel], m.head2.weight, t2[sel])
        loss = l1 + 0.5 * l2 + 1e-3 * bal
    loss.backward()
    torch.nn.utils.clip_grad_norm_(muon_p + dense_p, 1.0)      # P6: pool excluded
    opt.step(); sparse_sgd(a.poollr * f)
    if step % int(os.environ.get('MUON_EVERY', '2')) == 0: muon.step()
    with torch.no_grad():
        leaf = getattr(m, '_leaf', None)
        if leaf is not None and leaf.grad is not None:
            flat = rows.reshape(-1); g = leaf.grad
            gn = g.norm(dim=-1, keepdim=True).clamp_min(1e-8); qstep = m.codebook.detach().std()
            rstep = a.rowstep * qstep * f * (m.d ** 0.5); newrows = leaf.detach() - rstep * (g / gn)
            uid, inv = torch.unique(flat, return_inverse=True)
            acc = torch.zeros(uid.size(0), m.d, device=newrows.device, dtype=newrows.dtype); acc.index_copy_(0, inv, newrows)
            if not hasattr(m, '_pu'): m._pu = []; m._pa = []
            m._pu.append(uid); m._pa.append(acc)
            if step % int(os.environ.get('MUT_EVERY', '16')) == 0:
                pu = torch.cat(m._pu); pa = torch.cat(m._pa)
                u2, inv2 = torch.unique(pu, return_inverse=True)
                a2 = torch.zeros(u2.size(0), m.d, device=pa.device, dtype=pa.dtype); a2.index_copy_(0, inv2, pa)
                before = m.codons[u2].clone(); m.mutate_rows(u2, a2)
                m._mut = (m.codons[u2] != before).float().mean().item(); m._pu = []; m._pa = []
            m._leaf = None
    seen += a.batch * (a.seq - 1); toks = step * a.batch * (a.seq - 1)
    if step % a.log_every == 0:
        torch.cuda.synchronize(); dt = time.time() - t0; ts = seen / dt
        ts_hist.append(ts); avg = sum(ts_hist) / len(ts_hist)
        print(f'step={step} tokens={toks} ce={l1.item():.4f} ppl={math.exp(min(20, l1.item())):.1f} '
              f'bal={bal.item():.3f} mut={getattr(m,"_mut",0):.3f} tok_s={ts:,.0f} avg_tok_s={avg:,.0f}', flush=True)
        log_tp({'step': step, 'tokens': toks, 'ce': l1.item(), 'tok_s': ts, 'avg_tok_s': avg,
                'elapsed_h': (time.time() - wall0) / 3600, 'variant': 'v21'})
        t0 = time.time(); seen = 0
    if step % a.ckpt_every == 0:
        save(step); vl = valid_loss()
        print(f'CKPT step={step} valid_ce={vl:.4f} valid_ppl={math.exp(min(20, vl)):.1f}', flush=True)
        log_tp({'step': step, 'tokens': toks, 'valid_ce': vl,
                'avg_tok_s': sum(ts_hist)/max(1,len(ts_hist)), 'variant': 'v21'})
    if step % a.hf_every == 0: hf_push(step)
    if toks >= a.target_tokens:
        print(f'TARGET_TOKENS reached {toks}', flush=True); break
    if a.max_hours and (time.time() - wall0) / 3600 >= a.max_hours:
        print(f'MAX_HOURS reached {(time.time()-wall0)/3600:.2f}', flush=True); break
save(step); torch.save({'model': m.state_dict(), 'config': m.config()}, Path(a.out) / 'base.pt'); hf_push(step)
AVG = sum(ts_hist) / max(1, len(ts_hist))
print('PEER21_BASE_DONE', json.dumps({'valid_ce': valid_loss(), 'steps': step,
      'avg_tok_s': AVG, 'tokens': step * a.batch * (a.seq - 1)}), flush=True)