thefinalboss commited on
Commit
c20883a
·
verified ·
1 Parent(s): f27efd5

Upload scripts/fast4gpu_boost.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. scripts/fast4gpu_boost.py +210 -0
scripts/fast4gpu_boost.py ADDED
@@ -0,0 +1,210 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Fractus-1B boost trainer — B=4, compile, TF32, sequential scheduled sampling.
3
+
4
+ Resume from HF merge if pod weights are lost:
5
+ checkpoints/FRACTUS_1B_STAGE2_MERGED.pt
6
+
7
+ Usage (one process per GPU):
8
+ CUDA_VISIBLE_DEVICES=0 GPU_ID=0 python -u scripts/fast4gpu_boost.py
9
+ CUDA_VISIBLE_DEVICES=1 GPU_ID=1 python -u scripts/fast4gpu_boost.py
10
+ ...
11
+
12
+ Env:
13
+ GPU_ID, BATCH=4, SEQ=128, LR=7e-4, SS_RATE=0.25, SS_PROB=0.2
14
+ CKPT_IN — path to load (default: merged or per-gpu if present)
15
+ CKPT_OUT — path to save (default: checkpoints/fractus_1b_gpu{GPU}.pt)
16
+ START_TOKEN — optional integer resume offset into shard
17
+ SHARD — path to token shard .pt (int64 1D)
18
+ """
19
+ from __future__ import annotations
20
+
21
+ import os
22
+ import sys
23
+ import time
24
+ import json
25
+ import random
26
+ from pathlib import Path
27
+
28
+ import torch
29
+ import torch.nn.functional as F
30
+
31
+ ROOT = Path(__file__).resolve().parents[1]
32
+ sys.path.insert(0, str(ROOT))
33
+ os.chdir(ROOT)
34
+
35
+ from fractus.continuous_engine import ContinuousThoughtEngine
36
+
37
+ GPU = int(os.environ.get("GPU_ID", "0"))
38
+ LB_COEF = float(os.environ.get("LB_COEF", "0.02"))
39
+ GATE_TEMP = float(os.environ.get("GATE_TEMP", "2.5"))
40
+ LR = float(os.environ.get("LR", "7e-4"))
41
+ EMA_BETA = 0.98
42
+ SS_PROB = float(os.environ.get("SS_PROB", "0.2"))
43
+ SS_RATE = float(os.environ.get("SS_RATE", "0.25"))
44
+ B = int(os.environ.get("BATCH", "4"))
45
+ SEQ = int(os.environ.get("SEQ", "128"))
46
+
47
+ TARGET = dict(
48
+ d_model=1280,
49
+ n_heads=20,
50
+ d_head=64,
51
+ n_levels=2,
52
+ n_oscillators=16,
53
+ coupling_rank=8,
54
+ n_experts=128,
55
+ top_k=2,
56
+ expert_d_ff=2048,
57
+ siren_rank=64,
58
+ n_layers=16,
59
+ )
60
+
61
+ torch.manual_seed(42 + GPU)
62
+ torch.backends.cuda.matmul.allow_tf32 = True
63
+ torch.backends.cudnn.allow_tf32 = True
64
+ torch.backends.cudnn.benchmark = True
65
+ device = torch.device("cuda:0")
66
+
67
+ default_merged = ROOT / "checkpoints" / "FRACTUS_1B_STAGE2_MERGED.pt"
68
+ default_gpu = ROOT / "checkpoints" / f"fractus_1b_gpu{GPU}.pt"
69
+ CKPT_IN = Path(os.environ.get("CKPT_IN", str(default_gpu if default_gpu.exists() else default_merged)))
70
+ CKPT_OUT = Path(os.environ.get("CKPT_OUT", str(default_gpu)))
71
+ SHARD = Path(os.environ.get("SHARD", str(ROOT / "data" / f"shard_gpu{GPU}.pt")))
72
+
73
+ print(f"GPU {GPU}: BOOST B={B} SEQ={SEQ} LR={LR} SS_RATE={SS_RATE}", flush=True)
74
+ print(f"GPU {GPU}: load {CKPT_IN}", flush=True)
75
+
76
+ ck = torch.load(CKPT_IN, map_location="cpu", weights_only=False)
77
+ sd = ck.get("model_state", ck)
78
+ clean = {(k[10:] if k.startswith("_orig_mod.") else k): v for k, v in sd.items()}
79
+
80
+ eng = ContinuousThoughtEngine(vocab_size=50257, **{k: TARGET[k] for k in TARGET})
81
+ own = eng.state_dict()
82
+ loaded = 0
83
+ for k, v in clean.items():
84
+ if k in own and own[k].shape == v.shape:
85
+ own[k] = v
86
+ loaded += 1
87
+ elif (
88
+ k in own
89
+ and v.dim() >= 1
90
+ and own[k].dim() >= 1
91
+ and v.shape[0] > own[k].shape[0]
92
+ and v.shape[1:] == own[k].shape[1:]
93
+ ):
94
+ own[k] = v[: own[k].shape[0]].contiguous()
95
+ loaded += 1
96
+ eng.load_state_dict(own, strict=False)
97
+ print(f"GPU {GPU}: loaded_tensors={loaded}", flush=True)
98
+
99
+ with torch.no_grad():
100
+ for blk in eng.blocks:
101
+ if hasattr(blk, "moe") and hasattr(blk.moe, "temperature"):
102
+ blk.moe.temperature = GATE_TEMP
103
+
104
+ eng = eng.to(device)
105
+ eng.reset_thought(B)
106
+
107
+ try:
108
+ eng = torch.compile(eng, mode="reduce-overhead")
109
+ print(f"GPU {GPU}: compile OK", flush=True)
110
+ except Exception as e:
111
+ print(f"GPU {GPU}: compile skip: {e}", flush=True)
112
+
113
+ opt = torch.optim.SGD(eng.parameters(), lr=LR, momentum=0.9)
114
+
115
+ if not SHARD.exists():
116
+ raise FileNotFoundError(
117
+ f"Shard not found: {SHARD}\n"
118
+ "Place tokenized int64 1D shard at data/shard_gpu{id}.pt or set SHARD="
119
+ )
120
+ tokens = torch.load(SHARD, weights_only=False).to(torch.int64)
121
+ step_tokens = B * SEQ
122
+
123
+ start_token = int(os.environ.get("START_TOKEN", "0"))
124
+ # align to step
125
+ start_token = (start_token // step_tokens) * step_tokens
126
+ print(f"GPU {GPU}: RESUME start_token={start_token} step={step_tokens} shard_len={len(tokens)}", flush=True)
127
+
128
+ t0 = time.time()
129
+ ema_tf = None
130
+ ema_ss = None
131
+ n = 0
132
+ tok_sess = 0
133
+
134
+ CKPT_OUT.parent.mkdir(parents=True, exist_ok=True)
135
+
136
+ for start in range(start_token, len(tokens) - step_tokens - SEQ - 1, step_tokens):
137
+ chunk = tokens[start : start + step_tokens].view(B, SEQ).to(device)
138
+ target = tokens[start + 1 : start + step_tokens + 1].view(B, SEQ).to(device)
139
+
140
+ with torch.autocast("cuda", dtype=torch.bfloat16):
141
+ out = eng.tick_chunk_train(chunk)
142
+ logits, lb = out if isinstance(out, tuple) else (out, eng.last_lb_loss)
143
+ ce_tf = F.cross_entropy(logits.reshape(-1, logits.size(-1)), target.reshape(-1))
144
+ loss = ce_tf + LB_COEF * lb
145
+
146
+ opt.zero_grad(set_to_none=True)
147
+ loss.backward()
148
+ torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
149
+ opt.step()
150
+
151
+ tf_v = float(ce_tf.item())
152
+ lb_v = float(lb.detach().item()) if torch.is_tensor(lb) else float(lb)
153
+ ema_tf = tf_v if ema_tf is None else EMA_BETA * ema_tf + (1 - EMA_BETA) * tf_v
154
+
155
+ ce_ss_v = None
156
+ if random.random() < SS_RATE:
157
+ with torch.no_grad():
158
+ samp = torch.multinomial(
159
+ torch.softmax(logits.detach().float().reshape(-1, logits.size(-1)) / 0.9, dim=-1),
160
+ 1,
161
+ ).view(B, SEQ)
162
+ mixed = chunk.clone()
163
+ use_ss = torch.rand(B, SEQ, device=device) < SS_PROB
164
+ use_ss[:, 0] = False
165
+ prev = torch.cat([chunk[:, :1], samp[:, :-1]], dim=1)
166
+ mixed = torch.where(use_ss, prev, mixed)
167
+ with torch.autocast("cuda", dtype=torch.bfloat16):
168
+ out2 = eng.tick_chunk_train(mixed)
169
+ logits2, lb2 = out2 if isinstance(out2, tuple) else (out2, eng.last_lb_loss)
170
+ ce_ss = F.cross_entropy(logits2.reshape(-1, logits2.size(-1)), target.reshape(-1))
171
+ loss2 = 0.5 * ce_ss + LB_COEF * lb2
172
+ opt.zero_grad(set_to_none=True)
173
+ loss2.backward()
174
+ torch.nn.utils.clip_grad_norm_(eng.parameters(), 1.0)
175
+ opt.step()
176
+ ce_ss_v = float(ce_ss.item())
177
+ ema_ss = ce_ss_v if ema_ss is None else EMA_BETA * ema_ss + (1 - EMA_BETA) * ce_ss_v
178
+
179
+ n += 1
180
+ tok_sess += step_tokens
181
+
182
+ if n % 40 == 0:
183
+ tps = tok_sess / max(time.time() - t0, 1e-6)
184
+ extra = f" ss={ce_ss_v:.3f} ema_ss={ema_ss:.3f}" if ce_ss_v is not None and ema_ss is not None else ""
185
+ mem = torch.cuda.max_memory_allocated() / 1e9
186
+ print(
187
+ f"GPU {GPU}: {start + step_tokens:>12,} tf={tf_v:.3f} ema_tf={ema_tf:.3f}{extra} "
188
+ f"lb={lb_v:.3f} {tps:.0f} tok/s mem={mem:.1f}GB [boost]",
189
+ flush=True,
190
+ )
191
+
192
+ if n % 800 == 0:
193
+ torch.save(
194
+ {
195
+ "model_state": eng.state_dict(),
196
+ "config": {
197
+ **TARGET,
198
+ "gpu": GPU,
199
+ "boost": True,
200
+ "batch": B,
201
+ "lr": LR,
202
+ "ss_rate": SS_RATE,
203
+ "tokens_processed": start + step_tokens,
204
+ },
205
+ },
206
+ CKPT_OUT,
207
+ )
208
+ print(f"GPU {GPU}: saved [boost] -> {CKPT_OUT}", flush=True)
209
+
210
+ print(f"GPU {GPU}: DONE", flush=True)