| import argparse,torch | |
| from .tokenizer import ByteBPETokenizer | |
| from .model import AresTransformer | |
| from .config import AresConfig | |
| def main(): | |
| p=argparse.ArgumentParser();p.add_argument('--checkpoint',required=True);p.add_argument('--tokenizer',required=True);a=p.parse_args();t=ByteBPETokenizer.load(a.tokenizer);z=torch.load(a.checkpoint,map_location='cpu');m=AresTransformer(AresConfig(**z['config']));m.load_state_dict(z['model']);m.eval();print('Ares local CLI. Type /quit.') | |
| while (q:=input('you> '))!='/quit': | |
| ids=torch.tensor([t.encode('<user>'+q+'<assistant>',True)]);print('ares>',t.decode(m.generate(ids,80)[0].tolist())) | |
| if __name__=='__main__':main() | |