Zahid2005 commited on
Commit
3f2437f
·
verified ·
1 Parent(s): 8532440

Upload generate.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. generate.py +50 -0
generate.py ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import sys
3
+ import argparse
4
+ from gpt import GPT
5
+ from tokenization.character import vocab_size, decode, encode
6
+
7
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
8
+
9
+ def generate_text(prompt="Once upon a time", max_new_tokens=300, temperature=0.8, top_k=40, top_p=0.9):
10
+ model = GPT(
11
+ vocab_size=vocab_size,
12
+ d_model=512,
13
+ num_heads=8,
14
+ hidden_dim=2048,
15
+ num_layers=4,
16
+ attention_type="mha",
17
+ normalization_type="rms",
18
+ feedforward_type="swiglu",
19
+ position_encoding="sinusoidal"
20
+ ).to(device)
21
+
22
+ model.load_state_dict(torch.load("checkpoints/gpt_character.pth", map_location=device))
23
+ model.eval()
24
+
25
+ context = torch.tensor([encode(prompt)], dtype=torch.long, device=device)
26
+ generated = model.generate(context, max_new_tokens=max_new_tokens, temperature=temperature, top_k=top_k, top_p=top_p)
27
+ text = decode(generated[0].tolist())
28
+ return text
29
+
30
+ if __name__ == "__main__":
31
+ parser = argparse.ArgumentParser(description="Generate text from trained GPT model")
32
+ parser.add_argument("--prompt", type=str, default="Once upon a time", help="Initial prompt text")
33
+ parser.add_argument("--max_tokens", type=int, default=300, help="Number of tokens to generate")
34
+ parser.add_argument("--temp", type=float, default=0.8, help="Sampling temperature")
35
+ parser.add_argument("--top_k", type=int, default=40, help="Top-k filtering")
36
+ parser.add_argument("--top_p", type=float, default=0.9, help="Top-p (nucleus) filtering")
37
+ args = parser.parse_args()
38
+
39
+ print(f"\nPrompt: '{args.prompt}'")
40
+ print(f"Sampling Parameters: Temp={args.temp}, Top-k={args.top_k}, Top-p={args.top_p}")
41
+ print("=" * 60)
42
+ result = generate_text(
43
+ prompt=args.prompt,
44
+ max_new_tokens=args.max_tokens,
45
+ temperature=args.temp,
46
+ top_k=args.top_k,
47
+ top_p=args.top_p
48
+ )
49
+ print(result)
50
+ print("=" * 60)