eca / eca.py
jgalego's picture
Train 3000 steps
11821ba verified
Raw History Blame Contribute Delete
13.8 kB
# /// 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()