File size: 1,355 Bytes
3f431df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Simple causal completion CLI (no KV cache)."""
import argparse
from contextlib import nullcontext
import torch
from load_model import load_model

def main():
    p=argparse.ArgumentParser()
    p.add_argument('--model-dir',default='.')
    p.add_argument('--prompt',default='The scientific method is')
    p.add_argument('--max-new-tokens',type=int,default=64)
    p.add_argument('--temperature',type=float,default=0.0)
    p.add_argument('--device',default='cuda' if torch.cuda.is_available() else 'cpu')
    args=p.parse_args();torch.set_num_threads(4)
    model,tok=load_model(args.model_dir,args.device)
    ids=tok.encode(args.prompt).ids
    if not ids:raise ValueError('Prompt must encode to at least one token')
    eos=tok.token_to_id('<|endoftext|>')
    with torch.inference_mode():
        for _ in range(args.max_new_tokens):
            x=torch.tensor([ids[-model.block:]],device=args.device)
            ctx=torch.autocast('cuda',dtype=torch.bfloat16) if args.device.startswith('cuda') else nullcontext()
            with ctx:logits=model(x)[0][0,-1].float()
            nxt=int(logits.argmax()) if args.temperature<=0 else int(torch.multinomial(torch.softmax(logits/args.temperature,dim=-1),1))
            ids.append(nxt)
            if nxt==eos:break
    print(tok.decode(ids,skip_special_tokens=True))
if __name__=='__main__':main()