bind2_0 / release /src /train_bind2_0_babylm.py
SecludedCorner's picture
bind2_0: release package
8182d87 verified
Raw
History Blame Contribute Delete
4.14 kB
"""train_bind2_0_babylm.py -- train bind2_0 on the REAL BabyLM strict-small corpus (2026-07-12),
matched to bind1/mono's regime (SEQ256, batch, 150M tokens, flat-chunk of tokens_u16.bin).
CONTAMINATION-SAFE: each SEQ-block is a separate batch element, so the GDN state S starts from 0
per block (no cross-block carry) -- same data regime as the stateless bind1/mono flat-chunk runs.
Eval/inference forwards each probe independently (S from 0), so items never influence each other.
Usage: python train_bind2_0_babylm.py DIM DEPTH CHUNK SEQ BATCH TARGET TAG
Reads BABYLM_WORK/tokens_u16.bin (the real 16.3M-token BabyLM corpus). Writes {TAG}_final.pt +
train_{TAG}.json (arch bind2_0) + grid ckpts at BABYLM_CKPT_MARKS.
"""
import os, sys, time, json, random
import numpy as np
import torch, torch.nn.functional as F
from tokenizers import Tokenizer
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from modeling_bind2_0 import Bind2_0LM
WORK = os.environ["BABYLM_WORK"]
TOKJSON = os.path.join(WORK, "tokenizer.json")
TOKBIN = os.path.join(WORK, "tokens_u16.bin")
def main():
a = sys.argv
DIM, DEPTH, CHUNK = int(a[1]), int(a[2]), int(a[3])
SEQ, BATCH, TARGET, TAG = int(a[4]), int(a[5]), int(a[6]), a[7]
MLPH = int(a[8]) if len(a) > 8 else 576 # 832 => ~27.4M (BabyLM scale, matches mono); 576 => 23.9M (bind1 scale)
dev = "cuda"
SEED = int(os.environ.get("BABYLM_SEED", "0")); random.seed(SEED); torch.manual_seed(SEED); np.random.seed(SEED)
tok = Tokenizer.from_file(TOKJSON); V = tok.get_vocab_size()
toks = np.fromfile(TOKBIN, dtype=np.uint16).astype(np.int64); data = torch.from_numpy(toks)
nblk = (len(toks) - 1) // SEQ
assert nblk > 0
model = Bind2_0LM(V, DIM, DEPTH, 6, chunk=CHUNK, mlp_hidden=MLPH).to(dev).train()
P = sum(p.numel() for p in model.parameters())
print(f"bind2_0 params={P/1e6:.2f}M | REAL-BabyLM flat-chunk seq{SEQ}xb{BATCH} chunk{CHUNK} (S from 0 per block)", flush=True)
opt = torch.optim.AdamW(model.parameters(), lr=6e-4, betas=(0.9, 0.95), weight_decay=0.01)
g = torch.Generator().manual_seed(SEED)
torch.cuda.reset_peak_memory_stats(); t0 = time.time()
MARKS = sorted(int(float(m) * 1e6) for m in os.environ.get("BABYLM_CKPT_MARKS", "").split(",") if m.strip())
mi = 0; seen = 0; step = 0
while seen < TARGET:
idx = torch.randint(0, nblk, (BATCH,), generator=g)
x = torch.stack([data[i * SEQ:(i + 1) * SEQ + 1] for i in idx]).to(dev)
with torch.autocast("cuda", dtype=torch.bfloat16):
logits = model(x[:, :-1])
ce = F.cross_entropy(logits.reshape(-1, V).float(), x[:, 1:].reshape(-1))
opt.zero_grad(set_to_none=True); ce.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); opt.step()
step += 1; seen += BATCH * SEQ
while mi < len(MARKS) and seen >= MARKS[mi]:
os.makedirs(WORK + "/ckpt", exist_ok=True)
torch.save(model.state_dict(), WORK + f"/ckpt/{TAG}_g{MARKS[mi]/1e6:g}M.pt")
print(f" [ckpt-grid] {TAG}_g{MARKS[mi]/1e6:g}M.pt (seen={seen})", flush=True); mi += 1
if step == 1 or step % 200 == 0:
print(f"[{TAG}] st {step:6d} {seen/1e6:.1f}M ce {ce.item():.3f} ppl {np.exp(min(ce.item(),20)):.1f} | {round(seen/(time.time()-t0))} tok/s", flush=True)
dt = time.time() - t0
summ = {"tag": TAG, "arch": "bind2_0", "dim": DIM, "depth": DEPTH, "chunk": CHUNK, "mlp_hidden": MLPH,
"params_M": round(P / 1e6, 2), "vocab": V, "seq": SEQ, "batch": BATCH,
"peak_MiB": round(torch.cuda.max_memory_allocated() / 2**20), "tokens": seen,
"tok_s": round(seen / dt), "final_ce": round(ce.item(), 4), "ckpt_marks": os.environ.get("BABYLM_CKPT_MARKS", "")}
json.dump(summ, open(WORK + f"/train_{TAG}.json", "w", encoding="utf-8"), indent=2)
torch.save(model.state_dict(), WORK + f"/{TAG}_final.pt")
print(f"[{TAG}] DONE bind2_0 {P/1e6:.1f}M {seen/1e6:.1f}M ce {ce.item():.3f} ppl {np.exp(min(ce.item(),20)):.1f} {summ['tok_s']} tok/s peak {summ['peak_MiB']}MiB", flush=True)
if __name__ == "__main__":
main()