File size: 3,131 Bytes
758c573
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
import torch
import gradio as gr
from model import GPT, GPTConfig
import tiktoken
import os

# Configuration
device = 'cpu'
if torch.cuda.is_available():
    device = 'cuda'
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
    device = "mps"

print(f"Using device: {device}")

# Load Model
model_path = "final_model.pth"
if not os.path.exists(model_path):
    raise FileNotFoundError(f"Model file not found: {model_path}")

# Initialize model with same config as training
config = GPTConfig()
model = GPT(config)

# Load state dict
checkpoint = torch.load(model_path, map_location=device)
# Handle if checkpoint is a full checkpoint dict or just state_dict
if 'model_state_dict' in checkpoint:
    state_dict = checkpoint['model_state_dict']
else:
    state_dict = checkpoint

model.load_state_dict(state_dict)
model.to(device)
model.eval()

# Tokenizer
enc = tiktoken.get_encoding('gpt2')

def generate_text(start_text, max_new_tokens=100, temperature=0.8, top_k=50):
    if not start_text:
        return "Please enter some text to start."
    
    # Encode input
    start_ids = enc.encode(start_text)
    x = torch.tensor(start_ids, dtype=torch.long, device=device)[None, ...]
    
    # Generate
    with torch.no_grad():
        for _ in range(max_new_tokens):
            # Crop context if needed
            idx_cond = x if x.size(1) <= config.block_size else x[:, -config.block_size:]
            
            # Forward pass
            logits, _ = model(idx_cond)
            logits = logits[:, -1, :] / temperature
            
            # Top-k sampling
            if top_k is not None:
                v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
                logits[logits < v[:, [-1]]] = -float('Inf')
            
            probs = torch.nn.functional.softmax(logits, dim=-1)
            idx_next = torch.multinomial(probs, num_samples=1)
            
            x = torch.cat((x, idx_next), dim=1)
            
    # Decode output
    output_tokens = x[0].tolist()
    decoded_text = enc.decode(output_tokens)
    return decoded_text

# Gradio Interface
with gr.Blocks(title="GPT Text Generation") as demo:
    gr.Markdown("# GPT Text Generation")
    gr.Markdown("Enter some text and the model will continue it.")
    
    with gr.Row():
        with gr.Column():
            input_text = gr.Textbox(label="Input Text", placeholder="Once upon a time...", lines=5)
            with gr.Row():
                max_tokens = gr.Slider(minimum=10, maximum=500, value=100, step=10, label="Max New Tokens")
                temperature = gr.Slider(minimum=0.1, maximum=2.0, value=0.8, step=0.1, label="Temperature")
                top_k = gr.Slider(minimum=1, maximum=100, value=50, step=1, label="Top-K")
            generate_btn = gr.Button("Generate", variant="primary")
        
        with gr.Column():
            output_text = gr.Textbox(label="Generated Text", lines=10)
    
    generate_btn.click(
        fn=generate_text,
        inputs=[input_text, max_tokens, temperature, top_k],
        outputs=output_text
    )

if __name__ == "__main__":
    demo.launch()