"""The small models the course shows along the way (the "Sprout now" cards). python snapshots.py bigram # letter bigram: counts -> probabilities python snapshots.py mlp3 # 3 letters -> next letter (embeddings + hidden layer) python snapshots.py mlp8 # 8 letters -> next letter (two hidden layers) python snapshots.py emb2d # the same idea with 2-D embeddings, to draw them python snapshots.py mlptok # 8 BPE tokens -> next token Writes browser files into --out (weights in the same SPRT format as export.py). """ import argparse import json import math import os import time import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from export import pack # id 0 = start/end of a text, then '\n' and the 95 printable ASCII characters (as chars.js) ALPHABET = '\n' + ''.join(chr(i) for i in range(32, 127)) STOI = {c: i + 1 for i, c in enumerate(ALPHABET)} V = len(ALPHABET) + 1 def char_corpus(data, n_chars): """A slice of the pre-training text, as letters (decoded from the token file).""" from tokenizer import Tokenizer tok = Tokenizer.load(f'{data}/tokenizer.json') ids = np.memmap(f'{data}/train.bin', dtype=np.uint16, mode='r')[: n_chars // 3] eot = tok['<|endoftext|>'] text = tok.decode([int(i) for i in ids]).replace('<|endoftext|>', '\0') out = np.array([0 if c == '\0' else STOI.get(c, -1) for c in text], dtype=np.int64) return out[out >= 0] def bigram(args): ids = char_corpus(args.data, args.chars) counts = np.zeros((V, V), dtype=np.int64) np.add.at(counts, (ids[:-1], ids[1:]), 1) probs = (counts + 1) / (counts + 1).sum(1, keepdims=True) logits = np.log(probs).astype(np.float32) state = {'embed.weight': torch.from_numpy(logits)} cfg = {'vocab_size': V, 'context': 1} pack(state, cfg, {'kind': 'mlp', 'name': 'bigram-char', 'chars': int(len(ids))}, f'{args.out}/models/bigram-char.bin', quantize=False) with open(f'{args.out}/data/bigram-counts.json', 'w') as f: json.dump({'alphabet': ALPHABET, 'counts': counts.tolist(), 'chars': int(len(ids))}, f, separators=(',', ':')) nll = -np.log(probs[ids[:-1], ids[1:]]).mean() print(f'bigram: {len(ids)} letters, loss {nll:.4f} nats/letter') class MLP(nn.Module): def __init__(self, vocab, context, d_embed, hidden): super().__init__() self.context = context self.embed = nn.Embedding(vocab, d_embed) dims = [context * d_embed] + hidden self.hidden = nn.ModuleList(nn.Linear(a, b) for a, b in zip(dims, dims[1:])) self.out = nn.Linear(dims[-1], vocab) def forward(self, x): # x: (batch, context) h = self.embed(x).flatten(1) for layer in self.hidden: h = torch.tanh(layer(h)) return self.out(h) def windows(ids, context, pad=0): """(context letters, next letter) pairs; the window restarts at every text start.""" return torch.from_numpy(np.lib.stride_tricks.sliding_window_view(np.concatenate([np.full(context, pad), ids]), context + 1).copy()) def train_mlp(name, ids, vocab, context, d_embed, hidden, args, lr=3e-3, steps=None, batch=512): torch.manual_seed(0) data = windows(ids, context) n_val = len(data) // 50 val, tr = data[:n_val], data[n_val:] model = MLP(vocab, context, d_embed, hidden) opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.01) steps = steps or args.steps t0 = time.time() log = [] for step in range(steps + 1): if step % 500 == 0 or step == steps: with torch.no_grad(): idx = torch.randint(0, len(val), (8192,), generator=torch.Generator().manual_seed(0)) vb = val[idx] vl = F.cross_entropy(model(vb[:, :-1]), vb[:, -1]).item() log.append({'step': step, 'val': round(vl, 4)}) print(f'{name} step {step} val {vl:.4f} ({time.time() - t0:.0f}s)', flush=True) if step == steps: break b = tr[torch.randint(0, len(tr), (batch,))] loss = F.cross_entropy(model(b[:, :-1]), b[:, -1]) opt.zero_grad(set_to_none=True) loss.backward() opt.step() for g in opt.param_groups: g['lr'] = lr * (0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * step / steps))) return model, log def export_mlp(model, name, cfg, args, log): state = {k: v for k, v in model.state_dict().items()} size = pack(state, cfg, {'kind': 'mlp', 'name': name, 'log': log}, f'{args.out}/models/{name}.bin', quantize=True) print(f'{name}: {sum(p.numel() for p in model.parameters()) / 1e3:.0f}k params, {size / 1e3:.0f} KB') def main(): ap = argparse.ArgumentParser() ap.add_argument('what', choices=['bigram', 'mlp3', 'mlp8', 'emb2d', 'mlptok']) ap.add_argument('--data', default='data') ap.add_argument('--out', default='../../public/llm') ap.add_argument('--chars', type=int, default=30_000_000) ap.add_argument('--steps', type=int, default=20000) ap.add_argument('--threads', type=int, default=4) args = ap.parse_args() torch.set_num_threads(args.threads) os.makedirs(f'{args.out}/models', exist_ok=True) os.makedirs(f'{args.out}/data', exist_ok=True) if args.what == 'bigram': bigram(args) elif args.what in ('mlp3', 'mlp8', 'emb2d'): ids = char_corpus(args.data, args.chars) spec = {'mlp3': (3, 16, [256]), 'mlp8': (8, 24, [512, 512]), 'emb2d': (3, 2, [128])}[args.what] context, d_embed, hidden = spec name = {'mlp3': 'mlp-char-3', 'mlp8': 'mlp-char-8', 'emb2d': 'emb2d'}[args.what] model, log = train_mlp(name, ids, V, context, d_embed, hidden, args) if args.what == 'emb2d': emb = model.embed.weight.detach().numpy() with open(f'{args.out}/data/emb2d.json', 'w') as f: json.dump({'alphabet': ALPHABET, 'xy': emb.round(4).tolist(), 'log': log}, f, separators=(',', ':')) cfg = {'vocab_size': V, 'context': context, 'd_embed': d_embed, 'hidden': hidden} export_mlp(model, name, cfg, args, log) elif args.what == 'mlptok': ids = np.memmap(f'{args.data}/train.bin', dtype=np.uint16, mode='r')[:40_000_000].astype(np.int64) model, log = train_mlp('mlp-token', ids, 8192, 8, 48, [384], args, lr=2e-3, batch=512) cfg = {'vocab_size': 8192, 'context': 8, 'd_embed': 48, 'hidden': [384]} export_mlp(model, 'mlp-token', cfg, args, log) if __name__ == '__main__': main()