Download generate.py from SlayerLab/Slayer149: direct link, hf CLI and curl.
- Browser
- Download file 1.36 kB
-
https://huggingface.co/SlayerLab/Slayer149/resolve/main/generate.py
- Command line
-
hf download hf://SlayerLab/Slayer149/generate.py
-
curl -L -o generate.py https://huggingface.co/SlayerLab/Slayer149/resolve/main/generate.py
1.36 kB
| """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() | |