# /// script # requires-python = ">=3.12" # dependencies = [ # "huggingface_hub", # "jinja2", # "numpy", # "safetensors", # "tokenizers", # "torch", # ] # /// """Pre-train and sample Eça, a small GPT trained from scratch on the works of Eça de Queirós. Byte-level BPE tokenizer trained on the corpus, pre-norm decoder-only transformer with tied embeddings. Model quality is reported as bits per character on the last 2% of the text, so it stays comparable across tokenizers. """ # Torch modules only expose forward and keep their hyperparameters as attributes. # pylint: disable=too-few-public-methods,too-many-instance-attributes import argparse import copy import json import math import shutil import time from dataclasses import asdict, dataclass from pathlib import Path import torch import torch.nn.functional as F from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download, snapshot_download from huggingface_hub.errors import RepositoryNotFoundError from safetensors.torch import load_file, save_file from tokenizers import Tokenizer, decoders, models, pre_tokenizers, trainers from torch import nn REPO = "jgalego/eca" DATA = "jgalego/eca-queiros" DEVICE = "cuda" if torch.cuda.is_available() else "cpu" TOP_K = 40 HERE = Path(__file__).parent @dataclass class Config: """Architecture.""" vocab: int = 2048 block: int = 256 width: int = 256 layers: int = 6 heads: int = 4 dropout: float = 0.1 # Data def read_text(source): """Text of a local file or of eça.txt in the dataset repo.""" if Path(source).is_file(): return Path(source).read_text(encoding="utf-8") return Path(hf_hub_download(source, "eça.txt", repo_type="dataset")).read_text(encoding="utf-8") def new_tokenizer(text, vocab): """Train a byte-level BPE tokenizer on the corpus.""" tokenizer = Tokenizer(models.BPE()) tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False) tokenizer.decoder = decoders.ByteLevel() alphabet = pre_tokenizers.ByteLevel.alphabet() trainer = trainers.BpeTrainer(vocab_size=vocab, initial_alphabet=alphabet) tokenizer.train_from_iterator(text.splitlines(keepends=True), trainer) return tokenizer def prepare(args): """Tokenizer, training tokens, and the validation tokens with their length in characters. The first 98% of the text is for training, the rest for validation. """ text = read_text(args.data) tokenizer = new_tokenizer(text, args.vocab) cut = text.index("\n", int(len(text) * 0.98)) + 1 head, tail = (torch.tensor(tokenizer.encode(part).ids) for part in (text[:cut], text[cut:])) return tokenizer, head, (tail, len(text) - cut) # Model class Block(nn.Module): """Pre-norm transformer block with causal self-attention.""" def __init__(self, config): super().__init__() self.heads, self.dropout = config.heads, config.dropout self.norm1 = nn.LayerNorm(config.width) self.qkv = nn.Linear(config.width, 3 * config.width) self.proj = nn.Linear(config.width, config.width) self.norm2 = nn.LayerNorm(config.width) self.mlp = nn.Sequential( nn.Linear(config.width, 4 * config.width), nn.GELU(), nn.Linear(4 * config.width, config.width), ) self.drop = nn.Dropout(config.dropout) def forward(self, x): """Attention then MLP, each with a residual connection.""" batch, length, width = x.shape qkv = self.qkv(self.norm1(x)).view(batch, length, 3, self.heads, -1) q, k, v = qkv.permute(2, 0, 3, 1, 4) dropout = self.dropout if self.training else 0.0 att = F.scaled_dot_product_attention(q, k, v, is_causal=True, dropout_p=dropout) x = x + self.drop(self.proj(att.transpose(1, 2).reshape(batch, length, width))) return x + self.drop(self.mlp(self.norm2(x))) class Eca(nn.Module): """Decoder-only transformer. The output layer shares the token embedding.""" def __init__(self, config): super().__init__() self.config = config self.tokens = nn.Embedding(config.vocab, config.width) self.positions = nn.Embedding(config.block, config.width) self.drop = nn.Dropout(config.dropout) self.blocks = nn.Sequential(*(Block(config) for _ in range(config.layers))) self.norm = nn.LayerNorm(config.width) for embedding in (self.tokens, self.positions): nn.init.normal_(embedding.weight, std=0.02) def forward(self, ids): """Next-token logits (batch, length, vocab).""" x = self.tokens(ids) + self.positions(torch.arange(ids.shape[1], device=ids.device)) return self.norm(self.blocks(self.drop(x))) @ self.tokens.weight.T # Training def window_batch(data, args, generator): """Random windows of the token stream and the same windows shifted by one.""" block = args.block start = torch.randint(len(data) - block - 1, (args.batch_size,), generator=generator) x = torch.stack([data[i : i + block] for i in start]) y = torch.stack([data[i + 1 : i + block + 1] for i in start]) return x.to(DEVICE), y.to(DEVICE) @torch.no_grad() def validate(model, valid, batch_size=64): """Mean loss per token and bits per character over non-overlapping windows.""" valid, chars = valid model.eval() block = model.config.block count = (len(valid) - 1) // block x = valid[: count * block].view(count, block) y = valid[1 : count * block + 1].view(count, block) total = 0.0 for xs, ys in zip(x.split(batch_size), y.split(batch_size)): logits = model(xs.to(DEVICE)).float().reshape(-1, model.config.vocab) total += F.cross_entropy(logits, ys.to(DEVICE).reshape(-1), reduction="sum").item() model.train() loss = total / (count * block) return loss, loss * len(valid) / chars / math.log(2) def update(model, optimizer, batch, lr): """Take one optimizer step at the given learning rate and return the loss.""" for group in optimizer.param_groups: group["lr"] = lr x, y = batch with torch.autocast(DEVICE, dtype=torch.bfloat16, enabled=DEVICE == "cuda"): logits = model(x) loss = F.cross_entropy(logits.float().reshape(-1, model.config.vocab), y.reshape(-1)) optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() return loss.item() def schedule(args, step): """Linear warmup then cosine decay to zero.""" return args.lr * min(1, step / args.warmup) * 0.5 * (1 + math.cos(math.pi * step / args.steps)) def train(args): """Pre-train and keep the weights with the best validation score.""" tokenizer, train_ids, valid = prepare(args) torch.manual_seed(args.seed) generator = torch.Generator().manual_seed(args.seed) model = Eca(Config(args.vocab, args.block, args.width, args.layers, args.heads, args.dropout)) model.to(DEVICE) optimizer = torch.optim.AdamW(model.parameters(), args.lr, betas=(0.9, 0.95), weight_decay=0.1) best = {"val_bpc": math.inf, "weights": copy.deepcopy(model.state_dict())} start = time.time() for step in range(1, args.steps + 1): data = window_batch(train_ids, args, generator) loss = update(model, optimizer, data, schedule(args, step)) if step % args.eval_every == 0 or step == args.steps: val_loss, val_bpc = validate(model, valid) print(f"step {step} train {loss:.3f} val {val_loss:.3f} bpc {val_bpc:.3f}", flush=True) if val_bpc < best["val_bpc"]: best = { "step": step, "val_loss": round(val_loss, 4), "val_bpc": round(val_bpc, 4), "weights": copy.deepcopy(model.state_dict()), } model.load_state_dict(best.pop("weights")) training = { **best, "params": sum(p.numel() for p in model.parameters()), "steps": args.steps, "batch_size": args.batch_size, "lr": args.lr, "train_tokens": len(train_ids), "valid_tokens": len(valid[0]), "valid_chars": valid[1], "device": torch.cuda.get_device_name() if DEVICE == "cuda" else "cpu", "hours": round((time.time() - start) / 3600, 2), } save(model, tokenizer, Path(args.output), training) if args.push: push(Path(args.output), args.repo, f"Train {args.steps} steps") # Inference def resolve(name): """Local folder for a Hub repo id or a local path.""" return Path(name) if Path(name).is_dir() else Path(snapshot_download(name)) def load(name=REPO, device=DEVICE): """Load a trained model and its tokenizer from a Hub repo id or a local folder.""" folder = resolve(name) config = Config(**json.loads((folder / "config.json").read_text(encoding="utf-8"))) model = Eca(config) model.load_state_dict(load_file(folder / "model.safetensors")) return model.to(device).eval(), Tokenizer.from_file(str(folder / "tokenizer.json")) @torch.no_grad() def write(model, tokenizer, prompt, count=200, temperature=0.8): """Continue the prompt by sampling `count` tokens.""" ids = torch.tensor([tokenizer.encode(prompt).ids], device=next(model.parameters()).device) for _ in range(count): logits = model(ids[:, -model.config.block :])[:, -1] / temperature logits = logits.masked_fill(logits < logits.topk(TOP_K).values[:, -1:], -torch.inf) ids = torch.cat([ids, torch.multinomial(logits.softmax(-1), 1)], 1) return tokenizer.decode(ids[0].tolist()) def save(model, tokenizer, folder, training): """Write weights, config, tokenizer, code and training results to a folder.""" (folder / "results").mkdir(parents=True, exist_ok=True) state = {key: value.cpu().contiguous() for key, value in model.state_dict().items()} save_file(state, folder / "model.safetensors") tokenizer.save(str(folder / "tokenizer.json")) files = {"config.json": asdict(model.config), "results/train.json": training} for name, value in files.items(): (folder / name).write_text(json.dumps(value, indent=1) + "\n", encoding="utf-8") if Path(__file__).resolve() != (folder / "eca.py").resolve(): shutil.copy(__file__, folder / "eca.py") def push(folder, repo, message): """Upload a folder to a private model repo, creating it if needed.""" api = HfApi() api.create_repo(repo, private=True, exist_ok=True) api.upload_folder(folder_path=folder, repo_id=repo, commit_message=message) def card(args): """Render card.jinja into card/README.md with the config and results stored in the repo.""" files = {} try: folder = Path(snapshot_download(args.repo, allow_patterns=["config.json", "results/*"])) files = { path.relative_to(folder).as_posix(): json.loads(path.read_text(encoding="utf-8")) for path in folder.glob("**/*.json") } except RepositoryNotFoundError: pass data = ModelCardData( model_name=args.repo.split("/")[1], language="pt", license="mit", library_name="pytorch", pipeline_tag="text-generation", datasets=[DATA], tags=["portuguese", "literature", "from-scratch"], ) rendered = ModelCard.from_template( data, template_path=HERE / "card.jinja", repo=args.repo, data=DATA, config=files.get("config.json"), train=files.get("results/train.json"), ) (HERE / "card").mkdir(exist_ok=True) rendered.save(HERE / "card" / "README.md") def sample(args): """Print a continuation of the prompt.""" model, tokenizer = load(args.model) print(write(model, tokenizer, args.prompt, args.tokens, args.temperature)) def main(): """Parse arguments and run a command.""" parser = argparse.ArgumentParser(description=__doc__) commands = parser.add_subparsers(dest="command", required=True) train_parser = commands.add_parser("train") train_parser.add_argument("--data", default=DATA, help="dataset repo id or a local text file") train_parser.add_argument("--vocab", type=int, default=2048) train_parser.add_argument("--block", type=int, default=256) train_parser.add_argument("--width", type=int, default=256) train_parser.add_argument("--layers", type=int, default=6) train_parser.add_argument("--heads", type=int, default=4) train_parser.add_argument("--dropout", type=float, default=0.1) train_parser.add_argument("--steps", type=int, default=3000) train_parser.add_argument("--batch-size", type=int, default=64) train_parser.add_argument("--lr", type=float, default=1e-3) train_parser.add_argument("--warmup", type=int, default=200) train_parser.add_argument("--eval-every", type=int, default=250) train_parser.add_argument("--seed", type=int, default=0) train_parser.add_argument("--output", default="out/eca") train_parser.add_argument("--repo", default=REPO) train_parser.add_argument("--push", action="store_true") train_parser.set_defaults(run=train) write_parser = commands.add_parser("write") write_parser.add_argument("prompt") write_parser.add_argument("--model", default=REPO) write_parser.add_argument("--tokens", type=int, default=200) write_parser.add_argument("--temperature", type=float, default=0.8) write_parser.set_defaults(run=sample) card_parser = commands.add_parser("card") card_parser.add_argument("--repo", default=REPO) card_parser.set_defaults(run=card) args = parser.parse_args() args.run(args) if __name__ == "__main__": main()