thefinalboss commited on
Commit
a25e1b3
·
verified ·
1 Parent(s): bac0aaf

Upload scripts/fast4gpu_boost_v4.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. scripts/fast4gpu_boost_v4.py +333 -0
scripts/fast4gpu_boost_v4.py ADDED
@@ -0,0 +1,333 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Fractus-1B boost trainer v4 — AR gap.
3
+
4
+ v3 kernels/data/ckpt semantics, PLUS the losses that actually target free-run:
5
+
6
+ - P0 routing on (phase carry + per-token MoE + switch LB)
7
+ - SS_RATE default 1.0 (every step), SS_PROB ramps 0.2 → 0.5
8
+ - anti-repeat: λ · mean log p(class = last input token)
9
+ - unique@40 greedy probe (no ban) on a timer — THIS is the go/no-go, not ema_tf
10
+
11
+ Env (v3 plus):
12
+ SS_RATE=1.0 SS_PROB_START=0.2 SS_PROB_END=0.5 SS_RAMP_TOKENS=50000000
13
+ REPEAT_COEF=0.1 PROBE_EVERY=200
14
+ P0=1 (set 0 to ablate routing surgery)
15
+
16
+ One GPU:
17
+ CUDA_VISIBLE_DEVICES=$i GPU_ID=$i START_TOKEN=<manifest> \\
18
+ BATCH=8 CE_CHUNK=2048 FRACTUS_ATTN_IMPL=chunked BLOCK_CKPT=1 \\
19
+ python -u scripts/fast4gpu_boost_v4.py
20
+
21
+ Smoke ONE gpu 1–2h and read unique@40 before touching the other seven.
22
+ """
23
+ from __future__ import annotations
24
+
25
+ import os
26
+ import sys
27
+ import time
28
+ import json
29
+ from pathlib import Path
30
+
31
+ import torch
32
+
33
+ ROOT = Path(__file__).resolve().parents[1]
34
+ sys.path.insert(0, str(ROOT))
35
+ os.chdir(ROOT)
36
+
37
+ os.environ.setdefault("FRACTUS_ATTN_IMPL", os.environ.get("FRACTUS_ATTN_IMPL", "cumsum"))
38
+ from fractus.continuous_engine import ContinuousThoughtEngine
39
+ from fractus.generate_aligned import unique40_probe
40
+ from fractus.train.ar_loss import ss_prob_at
41
+ from fractus.train.v4_step import v4_forward_losses, v4_ss_pass, should_ss, snapshot_carry, restore_carry
42
+
43
+ GPU = int(os.environ.get("GPU_ID", "0"))
44
+ LB_COEF = float(os.environ.get("LB_COEF", "0.02"))
45
+ GATE_TEMP = float(os.environ.get("GATE_TEMP", "2.5"))
46
+ LR = float(os.environ.get("LR", "7e-4"))
47
+ EMA_BETA = float(os.environ.get("EMA_BETA", "0.98"))
48
+ SS_RATE = float(os.environ.get("SS_RATE", "1.0"))
49
+ SS_PROB_START = float(os.environ.get("SS_PROB_START", "0.2"))
50
+ SS_PROB_END = float(os.environ.get("SS_PROB_END", "0.5"))
51
+ SS_RAMP_TOKENS = int(os.environ.get("SS_RAMP_TOKENS", "50000000"))
52
+ REPEAT_COEF = float(os.environ.get("REPEAT_COEF", "0.1"))
53
+ PROBE_EVERY = int(os.environ.get("PROBE_EVERY", "200"))
54
+ P0 = os.environ.get("P0", "1") == "1"
55
+ B = int(os.environ.get("BATCH", "4"))
56
+ SEQ = int(os.environ.get("SEQ", "128"))
57
+ CE_CHUNK = int(os.environ.get("CE_CHUNK", "2048"))
58
+ ACCUM = max(1, int(os.environ.get("ACCUM", "1")))
59
+ BLOCK_CKPT = os.environ.get("BLOCK_CKPT", "0") == "1"
60
+ USE_COMPILE = os.environ.get("COMPILE", "0") == "1"
61
+ ATTN_IMPL = os.environ.get("FRACTUS_ATTN_IMPL", "cumsum")
62
+
63
+ TARGET = dict(
64
+ d_model=1280, n_heads=20, d_head=64, n_levels=2,
65
+ n_oscillators=16, coupling_rank=8, n_experts=128, top_k=2,
66
+ expert_d_ff=2048, siren_rank=64, n_layers=16,
67
+ )
68
+
69
+ torch.manual_seed(42 + GPU)
70
+ if torch.cuda.is_available():
71
+ torch.backends.cuda.matmul.allow_tf32 = True
72
+ torch.backends.cudnn.allow_tf32 = True
73
+ torch.backends.cudnn.benchmark = True
74
+ device = torch.device("cuda:0")
75
+ autocast = lambda: torch.autocast("cuda", dtype=torch.bfloat16)
76
+ else:
77
+ device = torch.device("cpu")
78
+ from contextlib import nullcontext
79
+ autocast = nullcontext
80
+
81
+ default_merged = ROOT / "checkpoints" / "FRACTUS_1B_STAGE2_MERGED.pt"
82
+ default_gpu = ROOT / "checkpoints" / f"fractus_1b_gpu{GPU}.pt"
83
+ CKPT_IN = Path(os.environ.get("CKPT_IN", str(default_gpu if default_gpu.exists() else default_merged)))
84
+ CKPT_OUT = Path(os.environ.get("CKPT_OUT", str(default_gpu)))
85
+ SHARD = Path(os.environ.get("SHARD", str(ROOT / "data" / f"shard_gpu{GPU}.npy")))
86
+ MANIFEST_OUT = Path(os.environ.get(
87
+ "MANIFEST_OUT", str(CKPT_OUT.parent / f"RESUME_MANIFEST_gpu{GPU}.json")))
88
+
89
+ print(
90
+ f"GPU {GPU}: BOOSTv4 B={B} SEQ={SEQ} LR={LR} SS_RATE={SS_RATE} "
91
+ f"ss_prob={SS_PROB_START}->{SS_PROB_END} repeat={REPEAT_COEF} P0={P0} "
92
+ f"attn={ATTN_IMPL} ce_chunk={CE_CHUNK} block_ckpt={BLOCK_CKPT}",
93
+ flush=True,
94
+ )
95
+ print(f"GPU {GPU}: load {CKPT_IN}", flush=True)
96
+
97
+ ck = torch.load(CKPT_IN, map_location="cpu", weights_only=False)
98
+ sd = ck.get("model_state", ck)
99
+ clean = {(k[10:] if k.startswith("_orig_mod.") else k): v for k, v in sd.items()}
100
+
101
+ eng = ContinuousThoughtEngine(vocab_size=50257, **TARGET)
102
+ own = eng.state_dict()
103
+ loaded = 0
104
+ for k, v in clean.items():
105
+ if k in own and own[k].shape == v.shape:
106
+ own[k] = v
107
+ loaded += 1
108
+ elif (
109
+ k in own and v.dim() >= 1 and own[k].dim() >= 1
110
+ and v.shape[0] > own[k].shape[0] and v.shape[1:] == own[k].shape[1:]
111
+ ):
112
+ own[k] = v[: own[k].shape[0]].contiguous()
113
+ loaded += 1
114
+ eng.load_state_dict(own, strict=False)
115
+ print(f"GPU {GPU}: loaded_tensors={loaded}", flush=True)
116
+
117
+ eng.set_p0_routing(
118
+ P0, lb_mode=("switch_topk" if P0 else "soft_var"),
119
+ )
120
+ with torch.no_grad():
121
+ for blk in eng.blocks:
122
+ if hasattr(blk, "moe") and hasattr(blk.moe, "temperature"):
123
+ blk.moe.temperature = GATE_TEMP
124
+
125
+ eng = eng.to(device)
126
+ eng.reset_thought(B)
127
+
128
+ if USE_COMPILE:
129
+ try:
130
+ eng = torch.compile(eng)
131
+ print(f"GPU {GPU}: torch.compile ON", flush=True)
132
+ except Exception as e:
133
+ print(f"GPU {GPU}: compile skip: {e}", flush=True)
134
+
135
+ payload = lambda: (eng._orig_mod if hasattr(eng, "_orig_mod") else eng)
136
+
137
+ opt = torch.optim.SGD(eng.parameters(), lr=LR, momentum=0.9)
138
+
139
+ if not SHARD.exists():
140
+ raise FileNotFoundError(f"Shard not found: {SHARD}")
141
+
142
+ import numpy as np
143
+ if not str(SHARD).endswith(".npy"):
144
+ raise FileNotFoundError(f"v4 expects .npy int32 shards, got {SHARD}")
145
+ shard_mm = np.load(str(SHARD), mmap_mode="r")
146
+ shard_len = int(shard_mm.shape[0])
147
+ print(f"GPU {GPU}: memmap shard {SHARD} len={shard_len:,} dtype={shard_mm.dtype}", flush=True)
148
+
149
+ step_tokens = B * SEQ
150
+
151
+
152
+ def fetch(start: int, count: int) -> torch.Tensor:
153
+ view = np.asarray(shard_mm[start : start + count])
154
+ return torch.from_numpy(view).to(torch.int64, non_blocking=True).to(device)
155
+
156
+
157
+ start_token = int(os.environ.get("START_TOKEN", "0"))
158
+ start_token = (start_token // step_tokens) * step_tokens
159
+ print(f"GPU {GPU}: RESUME start_token={start_token} step={step_tokens} shard_len={shard_len:,}",
160
+ flush=True)
161
+
162
+ t0 = time.time()
163
+ ema_tf = ema_ss = ema_rep = None
164
+ n = 0
165
+ tok_sess = 0
166
+ pending_backward = False
167
+ CKPT_OUT.parent.mkdir(parents=True, exist_ok=True)
168
+
169
+
170
+ def save_ckpt(tokens_done: int):
171
+ tmp = CKPT_OUT.with_suffix(CKPT_OUT.suffix + ".tmp")
172
+ torch.save(
173
+ {
174
+ "model_state": payload().state_dict(),
175
+ "config": {
176
+ **TARGET, "gpu": GPU, "boost_v4": True, "batch": B, "lr": LR,
177
+ "ss_rate": SS_RATE, "repeat_coef": REPEAT_COEF, "p0": P0,
178
+ "tokens_processed": tokens_done,
179
+ },
180
+ },
181
+ tmp,
182
+ )
183
+ os.replace(tmp, CKPT_OUT)
184
+ mtmp = MANIFEST_OUT.with_suffix(MANIFEST_OUT.suffix + ".tmp")
185
+ mtmp.write_text(json.dumps({
186
+ "ts": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
187
+ "gpu": GPU, "trainer": "fast4gpu_boost_v4",
188
+ "attn_impl": ATTN_IMPL, "p0": P0, "batch": B,
189
+ "tokens_processed": tokens_done, "start_token_next": tokens_done,
190
+ "shard": str(SHARD), "shard_len": shard_len, "ckpt": CKPT_OUT.name,
191
+ }, indent=1))
192
+ os.replace(mtmp, MANIFEST_OUT)
193
+ print(f"GPU {GPU}: saved [boostv4] -> {CKPT_OUT} @ {tokens_done:,} tok", flush=True)
194
+
195
+
196
+ def run_probe(tokens_done: int):
197
+ eng_p = payload()
198
+ thought = eng_p.thought_state
199
+ carries = [(b.attn_S.clone(), b.attn_z.clone(), b.kuramoto_phases.clone())
200
+ for b in eng_p.blocks]
201
+ B_live = thought.shape[0]
202
+ probe = unique40_probe(eng_p, max_new=40, mode="prefix")
203
+ probe_c = unique40_probe(eng_p, max_new=40, mode="carry")
204
+ # restore live train state
205
+ eng_p.thought_state = thought
206
+ for b, (S, z, ph) in zip(eng_p.blocks, carries):
207
+ b.attn_S, b.attn_z, b.kuramoto_phases = S, z, ph
208
+ eng_p.reset_thought(B_live) # batch may have been set to 1
209
+ # reset_thought zeros — put live carries back
210
+ eng_p.thought_state = thought
211
+ for b, (S, z, ph) in zip(eng_p.blocks, carries):
212
+ b.attn_S, b.attn_z, b.kuramoto_phases = S, z, ph
213
+ gate = "GO" if probe["gate_go"] else "NO-GO"
214
+ print(
215
+ f"GPU {GPU}: UNIQUE@40 PREFIX {gate} mean_u={probe['mean_unique']:.1f} "
216
+ f"echo={probe['mean_echo_frac']:.2f} | "
217
+ f"CARRY u={probe_c['mean_unique']:.1f} echo={probe_c['mean_echo_frac']:.2f} "
218
+ f"@ {tokens_done:,} tok",
219
+ flush=True,
220
+ )
221
+ for r in probe["rows"]:
222
+ print(
223
+ f" {r['prompt']!r}: unique={r['unique']} echo={r['echo_frac']:.2f} head={r['head']}",
224
+ flush=True,
225
+ )
226
+ rs = eng_p.routing_stats()
227
+ print(
228
+ f" routing alive={rs.get('alive_experts')} "
229
+ f"H={rs.get('dispatch_entropy', float('nan')):.3f} "
230
+ f"max_frac={rs.get('max_frac', float('nan')):.3f}",
231
+ flush=True,
232
+ )
233
+ return probe
234
+
235
+
236
+ for start in range(start_token, shard_len - step_tokens - SEQ - 1, step_tokens):
237
+ block = fetch(start, step_tokens + 1)
238
+ chunk = block[:step_tokens].view(B, SEQ).long()
239
+ target = block[1:].view(B, SEQ)
240
+ tokens_now = start + step_tokens
241
+ ss_prob = ss_prob_at(tokens_now, SS_PROB_START, SS_PROB_END, SS_RAMP_TOKENS)
242
+ carry_snap = snapshot_carry(eng)
243
+
244
+ with autocast():
245
+ loss, extras = v4_forward_losses(
246
+ payload(), chunk, target,
247
+ lb_coef=LB_COEF, repeat_coef=REPEAT_COEF,
248
+ ce_chunk=CE_CHUNK, block_ckpt=BLOCK_CKPT,
249
+ )
250
+ ce_tf, lb, rep, h = extras["ce"], extras["lb"], extras["repeat"], extras["h"]
251
+
252
+ ss_fired = False
253
+ ce_ss_v = None
254
+ if should_ss(SS_RATE):
255
+ ss_fired = True
256
+
257
+ if ACCUM == 1:
258
+ loss.backward()
259
+ torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
260
+ opt.step()
261
+ opt.zero_grad(set_to_none=True)
262
+ if ss_fired:
263
+ restore_carry(eng, carry_snap)
264
+ with autocast():
265
+ loss2, ce_ss = v4_ss_pass(
266
+ payload(), chunk, target, h.detach(),
267
+ ss_prob=ss_prob, lb_coef=LB_COEF,
268
+ ce_chunk=CE_CHUNK, block_ckpt=False,
269
+ )
270
+ loss2.backward()
271
+ torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
272
+ opt.step()
273
+ opt.zero_grad(set_to_none=True)
274
+ ce_ss_v = float(ce_ss.item())
275
+ ema_ss = ce_ss_v if ema_ss is None else EMA_BETA * ema_ss + (1 - EMA_BETA) * ce_ss_v
276
+ else:
277
+ (loss / ACCUM).backward()
278
+ if ss_fired:
279
+ restore_carry(eng, carry_snap)
280
+ with autocast():
281
+ loss2, ce_ss = v4_ss_pass(
282
+ payload(), chunk, target, h.detach(),
283
+ ss_prob=ss_prob, lb_coef=LB_COEF,
284
+ ce_chunk=CE_CHUNK, block_ckpt=False,
285
+ )
286
+ (loss2 / ACCUM).backward()
287
+ ce_ss_v = float(ce_ss.item())
288
+ ema_ss = ce_ss_v if ema_ss is None else EMA_BETA * ema_ss + (1 - EMA_BETA) * ce_ss_v
289
+ pending_backward = True
290
+
291
+ tf_v = float(ce_tf.detach().item())
292
+ lb_v = float(lb.detach().item()) if torch.is_tensor(lb) else float(lb)
293
+ rp_v = float(rep.detach().item())
294
+ ema_tf = tf_v if ema_tf is None else EMA_BETA * ema_tf + (1 - EMA_BETA) * tf_v
295
+ ema_rep = rp_v if ema_rep is None else EMA_BETA * ema_rep + (1 - EMA_BETA) * rp_v
296
+
297
+ n += 1
298
+ tok_sess += step_tokens
299
+
300
+ if ACCUM > 1 and n % ACCUM == 0:
301
+ torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
302
+ opt.step()
303
+ opt.zero_grad(set_to_none=True)
304
+ pending_backward = False
305
+
306
+ if n % 40 == 0:
307
+ tps = tok_sess / max(time.time() - t0, 1e-6)
308
+ extra = f" ss={ce_ss_v:.3f} ema_ss={ema_ss:.3f}" if ce_ss_v is not None else ""
309
+ try:
310
+ mem_s = f" mem={torch.cuda.max_memory_allocated() / 1e9:.1f}GB"
311
+ except Exception:
312
+ mem_s = ""
313
+ print(
314
+ f"GPU {GPU}: {tokens_now:>12,} tf={tf_v:.3f} ema_tf={ema_tf:.3f}{extra} "
315
+ f"rep={rp_v:.3f} ema_rep={ema_rep:.3f} lb={lb_v:.3f} "
316
+ f"ssp={ss_prob:.2f} {tps:.0f} tok/s{mem_s} [boostv4]",
317
+ flush=True,
318
+ )
319
+
320
+ if PROBE_EVERY > 0 and n % PROBE_EVERY == 0:
321
+ run_probe(tokens_now)
322
+
323
+ # All 8 GPUs save, staggered so they never write 8x4.3G at once.
324
+ if n % 4000 == (GPU * 500) % 4000:
325
+ save_ckpt(tokens_now)
326
+
327
+ if pending_backward:
328
+ torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
329
+ opt.step()
330
+ opt.zero_grad(set_to_none=True)
331
+
332
+ save_ckpt(start_token + n * step_tokens)
333
+ print(f"GPU {GPU}: DONE", flush=True)