Download train_seq.py from basically-experimental/Notio-3.7M-RNN-v1: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/basically-experimental/Notio-3.7M-RNN-v1/resolve/main/train_seq.py
- Command line
-
hf download hf://basically-experimental/Notio-3.7M-RNN-v1/train_seq.py
-
curl -L -o train_seq.py https://huggingface.co/basically-experimental/Notio-3.7M-RNN-v1/resolve/main/train_seq.py
10.3 kB
| #!/usr/bin/env python3 | |
| """Stateful truncated BPTT training loop for the RNN stack. | |
| - layer_0.seq_windows(): B contiguous tapes walked in order | |
| - RNN state carried across windows but detached at every boundary: | |
| gradients are truncated to one window, while information flows forward | |
| through the carried state (the k1 == k2 == T form of truncated BPTT) | |
| - state reset to zeros only at sweep boundaries (the one real discontinuity) | |
| Usage: | |
| python3 src/train_seq.py # train (defaults: 1 epoch, cpu/cuda auto) | |
| python3 src/train_seq.py --check # prove truncation + statefulness | |
| python3 src/train_seq.py --epochs 3 --lr 3e-4 --save out.pt | |
| """ | |
| import argparse | |
| import math | |
| import sys | |
| import time | |
| from itertools import islice | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| PROJ = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(PROJ)) | |
| sys.path.insert(0, str(PROJ / "scripts")) | |
| import tokenizer as tk | |
| from model import Notio, NotioConfig | |
| tk.load_vocab() # also asserts vocab.txt == vocab.bin at train start | |
| def detach_states(states): | |
| """Detach every hidden/cell tensor: no gradient flows across windows. | |
| Handles LSTM (h, c) tuples and GRU bare tensors.""" | |
| out = [] | |
| for st in states: | |
| out.append(tuple(s.detach() for s in st) if isinstance(st, tuple) else st.detach()) | |
| return out | |
| def ids_to_display(ids): | |
| """id tensor -> human text via the runtime decoder (tags dropped/mapped).""" | |
| out = bytes(ids.cpu().numpy().astype(np.uint8)).translate(tk.TABLE_DEC) | |
| for sent, tok in tk.SENT_DEC: | |
| out = out.replace(sent, tok) | |
| return tk.display(out).decode() | |
| def generate(m, device, max_tokens=256, temperature=0.9, top_k=0): | |
| """Sample one story: <bos> prompt, carry RNN state, stop at <eos>.""" | |
| m.eval() | |
| ids = torch.tensor([[1]], dtype=torch.long, device=device) # <bos> | |
| states = None | |
| pos = 0 | |
| for _ in range(max_tokens): | |
| logits, states = m(ids[:, -1:], states, pos_offset=pos) | |
| states = detach_states(states) | |
| pos += 1 | |
| logits = logits[:, -1, :] / max(temperature, 1e-6) | |
| if top_k > 0: | |
| v, _ = torch.topk(logits, top_k) | |
| logits = torch.where(logits < v[:, -1:], | |
| torch.full_like(logits, -float("inf")), logits) | |
| probs = F.softmax(logits, dim=-1) | |
| nxt = torch.multinomial(probs, 1) | |
| ids = torch.cat([ids, nxt], dim=1) | |
| if nxt.item() == 2: # <eos> | |
| break | |
| m.train() | |
| return ids | |
| def sample_and_print(m, device, n=2): | |
| for i in range(n): | |
| ids = generate(m, device) | |
| print(f"--- sample {i} ({ids.numel()} ids) ---") | |
| print(ids_to_display(ids)) | |
| def val_loss(m, device, n_blocks=20): | |
| """Mean next-token loss over the first n_blocks contiguous val blocks, | |
| carrying state across them (mirrors training).""" | |
| m.eval() | |
| states = None | |
| total = 0.0 | |
| n = min(n_blocks, m.layer0.n_val_blocks) | |
| for i in range(n): | |
| x, y = m.layer0.val_block(i) | |
| x, y = x.unsqueeze(0).to(device), y.unsqueeze(0).to(device) | |
| logits, states = m(x, states) | |
| total += F.cross_entropy(logits.view(-1, m.cfg.head.vocab_size), y.view(-1)).item() | |
| states = detach_states(states) | |
| m.train() | |
| return total / n | |
| def lr_at(step, lr, warmup, total, min_lr): | |
| """Linear warmup, then cosine decay to min_lr over [warmup, total).""" | |
| if step < warmup: | |
| return lr * step / max(warmup, 1) | |
| if step >= total: | |
| return min_lr | |
| prog = (step - warmup) / max(total - warmup, 1) | |
| return min_lr + 0.5 * (lr - min_lr) * (1 + math.cos(math.pi * prog)) | |
| def check(): | |
| torch.manual_seed(0) | |
| m = Notio(NotioConfig()) | |
| m.train() | |
| it = m.layer0.seq_windows() | |
| x1, y1 = next(it) | |
| x2, y2 = next(it) | |
| # 1) statefulness: carried state changes the second window's output | |
| _, st = m(x1) | |
| logits_carry, _ = m(x2, st) | |
| logits_zero, _ = m(x2) | |
| assert not torch.equal(logits_carry, logits_zero), "state is not used" | |
| # 2) truncation: detach cuts the graph back to window 1. | |
| # (values still flow - that is the point of carrying state) | |
| V = m.cfg.head.vocab_size | |
| assert st[0][0].grad_fn is not None, "state should carry history to window 1" | |
| st_d = detach_states(st) | |
| assert st_d[0][0].grad_fn is None, "detached state should be a leaf" | |
| logits_d, _ = m(x2, st_d) | |
| assert torch.equal(logits_carry, logits_d), "detach changed the forward pass" | |
| def pgrad(states_): | |
| m.zero_grad() | |
| logits, _ = m(x2, states_) | |
| F.cross_entropy(logits.view(-1, V), y2.view(-1)).backward() | |
| return sum(p.grad.detach().abs().sum() for p in m.parameters() if p.grad is not None) | |
| gd = pgrad(st_d) # detached: graph stops at the carried state | |
| gu = pgrad(st) # undetached: gradient also flows back through window 1 | |
| assert torch.equal(gd, gu) is False, "detach did not truncate the gradient path" | |
| # 3) tape continuity: window w+1 continues window w inside a tape | |
| assert torch.equal(x2[:, 0], y1[:, -1]), "tapes are not contiguous" | |
| # 4) a few real optimizer steps with carried+detached state | |
| opt = torch.optim.AdamW(m.parameters(), lr=1e-3) | |
| states = None | |
| it2 = m.layer0.seq_windows() | |
| t0 = time.time() | |
| for step in range(3): | |
| x, y = next(it2) | |
| logits, states = m(x, states) | |
| loss = F.cross_entropy(logits.view(-1, m.cfg.head.vocab_size), y.view(-1)) | |
| opt.zero_grad() | |
| loss.backward() | |
| opt.step() | |
| states = detach_states(states) | |
| print(f"step {step}: loss {loss.item():.3f}") | |
| toks = 3 * x.numel() | |
| print(f"truncation: True | statefulness: True | tape continuity: True") | |
| print(f"{toks / (time.time() - t0):,.0f} tokens/s (CPU, {m.n_params:,} params, " | |
| f"T={m.cfg.layer0.block_size}, B={m.cfg.layer0.batch_size}, d={m.cfg.layer1.d_model}, " | |
| f"n_blocks={m.cfg.n_blocks})") | |
| def main(device, epochs, lr, log_every, save, max_steps, save_every, | |
| warmup, min_lr, resume, val_blocks): | |
| torch.backends.cudnn.benchmark = True | |
| if resume: | |
| ck = torch.load(resume, map_location=device, weights_only=False) # our own file, contains NotioConfig | |
| m = Notio(ck["cfg"]) | |
| m.load_state_dict(ck["state_dict"]) | |
| m.to(device) | |
| start = ck.get("step", 0) | |
| print(f"resumed {resume} @ step {start} " | |
| f"(loss {ck.get('loss', float('nan')):.3f}, " | |
| f"val loss {ck.get('val_loss', float('nan')):.3f})") | |
| else: | |
| m = Notio(NotioConfig()) | |
| m.to(device) | |
| start = 0 | |
| m.train() | |
| opt = torch.optim.AdamW(m.parameters(), lr=lr) | |
| total = max_steps if max_steps else epochs * m.layer0.n_seq_windows | |
| states = None | |
| step = start | |
| t0 = time.time() | |
| it = islice(m.layer0.seq_windows(), start, None) # continue the sweep position | |
| for epoch in range(epochs if max_steps is None else 1): | |
| src = islice(it, max_steps) if max_steps is not None else it | |
| for x, y in src: | |
| new_lr = lr_at(step, lr, warmup, total, min_lr) | |
| for g in opt.param_groups: | |
| g["lr"] = new_lr | |
| x, y = x.to(device), y.to(device) | |
| logits, states = m(x, states) | |
| loss = F.cross_entropy(logits.view(-1, m.cfg.head.vocab_size), y.view(-1)) | |
| opt.zero_grad() | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(m.parameters(), 1.0) | |
| opt.step() | |
| states = detach_states(states) # truncate gradients at window length T | |
| step += 1 | |
| if step % log_every == 0: | |
| tok = step * x.numel() | |
| print(f"epoch {epoch} step {step}: loss {loss.item():.3f} | " | |
| f"{tok / (time.time() - t0):,.0f} tok/s | lr {new_lr:.2e}") | |
| if step % save_every == 0: | |
| vl = val_loss(m, device, val_blocks) | |
| print(f"val loss @ step {step}: {vl:.3f}") | |
| sample_and_print(m, device) | |
| if save: | |
| torch.save({"step": step, "loss": loss.item(), "val_loss": vl, | |
| "lr": new_lr, "window": step, "cfg": m.cfg, | |
| "state_dict": m.state_dict()}, save) | |
| print(f"checkpoint saved: {save} @ step {step}") | |
| states = None # sweep ended: reset at the real discontinuity | |
| if max_steps is not None: | |
| break | |
| it = m.layer0.seq_windows() | |
| total_tok = (step - start) * m.cfg.layer0.batch_size * m.cfg.layer0.block_size | |
| vl = val_loss(m, device, val_blocks) | |
| print(f"done: {step} steps | {total_tok / (time.time() - t0):,.0f} tok/s on {device} " | |
| f"| final loss {loss.item():.3f} | val loss {vl:.3f}") | |
| if save: | |
| torch.save({"step": step, "loss": loss.item(), "val_loss": vl, "lr": new_lr, | |
| "window": step, "cfg": m.cfg, "state_dict": m.state_dict()}, save) | |
| print(f"saved {save}") | |
| if __name__ == "__main__": | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--check", action="store_true", help="prove truncation + statefulness, then exit") | |
| ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") | |
| ap.add_argument("--epochs", type=int, default=1) | |
| ap.add_argument("--lr", type=float, default=4e-4) | |
| ap.add_argument("--log-every", type=int, default=20) | |
| ap.add_argument("--max-steps", type=int, default=None, help="benchmark: train this many steps then stop") | |
| ap.add_argument("--save", default=None, help="checkpoint path (saved every --save-every steps + at end)") | |
| ap.add_argument("--save-every", type=int, default=500) | |
| ap.add_argument("--warmup", type=int, default=500) | |
| ap.add_argument("--min-lr", type=float, default=4e-5) | |
| ap.add_argument("--resume", default=None, help="checkpoint to resume from (continues sweep position)") | |
| ap.add_argument("--val-blocks", type=int, default=20) | |
| a = ap.parse_args() | |
| if a.check: | |
| check() | |
| else: | |
| main(a.device, a.epochs, a.lr, a.log_every, a.save, a.max_steps, a.save_every, | |
| a.warmup, a.min_lr, a.resume, a.val_blocks) | |