File size: 1,046 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
import argparse
import sys
from pathlib import Path

import torch
torch.set_num_threads(1)

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

from krull import CharTokenizer, KRULLConfig, KRULLNano


def main():
    p = argparse.ArgumentParser()
    p.add_argument('--model', default='artifacts/krull_nano.pt')
    p.add_argument('--tokenizer', default='artifacts/tokenizer.json')
    p.add_argument('--prompt', default='KRULL is')
    p.add_argument('--max-new-tokens', type=int, default=120)
    p.add_argument('--device', default='cpu')
    args = p.parse_args()

    tok = CharTokenizer.load(args.tokenizer)
    ckpt = torch.load(args.model, map_location=args.device)
    cfg = KRULLConfig(**ckpt['config'])
    model = KRULLNano(cfg).to(args.device)
    model.load_state_dict(ckpt['model'])

    x = torch.tensor([tok.encode(args.prompt)], dtype=torch.long, device=args.device)
    y = model.generate(x, max_new_tokens=args.max_new_tokens)
    print(tok.decode(y[0].tolist()))


if __name__ == '__main__':
    main()