File size: 2,384 Bytes
b8daeef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
import argparse
import json
import sys
from pathlib import Path

import torch
torch.set_num_threads(1)
from torch.utils.data import Dataset, DataLoader

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

from krull import CharTokenizer, KRULLConfig, KRULLNano


class TextDataset(Dataset):
    def __init__(self, ids, block_size):
        self.ids = torch.tensor(ids, dtype=torch.long)
        self.block_size = block_size

    def __len__(self):
        return max(0, len(self.ids) - self.block_size - 1)

    def __getitem__(self, i):
        x = self.ids[i:i+self.block_size]
        y = self.ids[i+1:i+self.block_size+1]
        return x, y


def main():
    p = argparse.ArgumentParser()
    p.add_argument('--config', default='configs/krull_nano.json')
    p.add_argument('--tokenizer', default='artifacts/tokenizer.json')
    p.add_argument('--data', default='data/tiny_corpus.txt')
    p.add_argument('--out', default='artifacts/krull_nano.pt')
    p.add_argument('--epochs', type=int, default=5)
    p.add_argument('--batch-size', type=int, default=16)
    p.add_argument('--lr', type=float, default=3e-4)
    p.add_argument('--device', default='cpu')
    args = p.parse_args()

    tok = CharTokenizer.load(args.tokenizer)
    cfg_data = json.loads(Path(args.config).read_text(encoding='utf-8'))
    cfg = KRULLConfig(vocab_size=tok.vocab_size, **cfg_data)
    model = KRULLNano(cfg).to(args.device)

    text = Path(args.data).read_text(encoding='utf-8')
    ids = tok.encode(text)
    ds = TextDataset(ids, cfg.block_size)
    if len(ds) == 0:
        raise RuntimeError('Dataset is too small for the configured block_size.')
    dl = DataLoader(ds, batch_size=args.batch_size, shuffle=True)
    opt = torch.optim.AdamW(model.parameters(), lr=args.lr)

    model.train()
    for epoch in range(args.epochs):
        total = 0.0
        for x, y in dl:
            x, y = x.to(args.device), y.to(args.device)
            _, loss = model(x, y)
            opt.zero_grad()
            loss.backward()
            opt.step()
            total += loss.item()
        print(f'epoch {epoch+1}/{args.epochs} loss={total/len(dl):.4f}')

    Path(args.out).parent.mkdir(parents=True, exist_ok=True)
    torch.save({'config': cfg.__dict__, 'model': model.state_dict()}, args.out)
    print(f'Model saved to {args.out}')


if __name__ == '__main__':
    main()