Custom-GPT-40M-Base / example_generate.py
sraivante's picture
Publish Custom GPT 40M base model with training data
87842ae verified
Raw History Blame Contribute Delete
1.77 kB
"""Generate a short English continuation with Custom GPT 40M.
Copyright (c) 2026 sraivante. SPDX-License-Identifier: Apache-2.0
"""
import argparse
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
def main():
parser=argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model",default="sraivante/Custom-GPT-40M-Base")
parser.add_argument("--revision",default="main",help="Pin a Hub commit for reproducibility.")
parser.add_argument("--prompt",default="The future of artificial intelligence")
parser.add_argument("--max-new-tokens",type=int,default=80)
parser.add_argument("--seed",type=int,default=42)
parser.add_argument("--greedy",action="store_true")
parser.add_argument("--device",default="cpu",choices=["cpu","cuda"])
args=parser.parse_args()
if not args.prompt.strip():parser.error("Provide a non-empty English prompt.")
if args.max_new_tokens<1:parser.error("--max-new-tokens must be positive.")
tokenizer=AutoTokenizer.from_pretrained(args.model,revision=args.revision)
model=AutoModelForCausalLM.from_pretrained(args.model,revision=args.revision).to(args.device).eval()
inputs=tokenizer(args.prompt,return_tensors="pt",add_special_tokens=False).to(args.device)
remaining=model.config.n_positions-inputs["input_ids"].shape[1]
if remaining<1:parser.error("The prompt must be shorter than 256 tokens.")
torch.manual_seed(args.seed)
options={} if args.greedy else {"temperature":0.8,"top_k":50}
with torch.inference_mode():
output=model.generate(**inputs,max_new_tokens=min(args.max_new_tokens,remaining),do_sample=not args.greedy,**options)
print(tokenizer.decode(output[0],skip_special_tokens=True))
if __name__=="__main__":main()