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()
|