Views
No views yet
pip install transformers torch1import tkinter as tk
2from tkinter import scrolledtext, messagebox, ttk
3import threading
4import torch
5from transformers import AutoModelForCausalLM, AutoTokenizer
6
7# Model name
8MODEL_NAME = "gss1147/GPT5.1-high-reasoning-codex-0.4B"
9
10# Global variables for model and tokenizer (loaded once)
11model = None
12tokenizer = None
13
14def load_model():
15 """Load the model and tokenizer from Hugging Face."""
16 global model, tokenizer
17 try:
18 # Show loading message in the GUI (if root already exists)
19 if 'root' in globals():
20 status_label.config(text="Loading model... This may take a while.")
21 root.update()
22
23 device = "cuda" if torch.cuda.is_available() else "cpu"
24 tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
25 model = AutoModelForCausalLM.from_pretrained(
26 MODEL_NAME,
27 torch_dtype=torch.float16 if device == "cuda" else torch.float32,
28 device_map="auto" if device == "cuda" else None
29 )
30 if device == "cpu":
31 model.to(device)
32 model.eval()
33
34 if 'root' in globals():
35 status_label.config(text="Model loaded. Ready.")
36 generate_button.config(state=tk.NORMAL)
37 except Exception as e:
38 messagebox.showerror("Error", f"Failed to load model:\n{e}")
39 if 'root' in globals():
40 status_label.config(text="Loading failed.")
41 raise
42
43def generate_text():
44 """Start generation in a separate thread."""
45 thread = threading.Thread(target=generate_thread)
46 thread.start()
47
48def generate_thread():
49 """Run generation and update the GUI when done."""
50 # Disable the generate button and show progress
51 generate_button.config(state=tk.DISABLED, text="Generating...")
52 status_label.config(text="Generating...")
53 root.update()
54
55 try:
56 # Get parameters from the GUI
57 prompt = input_text.get("1.0", tk.END).strip()
58 if not prompt:
59 messagebox.showwarning("Warning", "Please enter a prompt.")
60 return
61
62 max_new = int(max_new_tokens_var.get())
63 temp = float(temperature_var.get())
64 top_p_val = float(top_p_var.get())
65 rep_penalty = float(rep_penalty_var.get())
66
67 # Tokenize input
68 inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
69
70 # Generate
71 with torch.no_grad():
72 outputs = model.generate(
73 **inputs,
74 max_new_tokens=max_new,
75 temperature=temp,
76 top_p=top_p_val,
77 repetition_penalty=rep_penalty,
78 do_sample=True,
79 pad_token_id=tokenizer.eos_token_id
80 )
81
82 # Decode and display
83 generated = tokenizer.decode(outputs[0], skip_special_tokens=True)
84 # Remove the input prompt from the output if desired
85 if generated.startswith(prompt):
86 generated = generated[len(prompt):].lstrip()
87
88 # Update output in GUI (thread-safe via after)
89 root.after(0, lambda: display_output(generated))
90
91 except Exception as e:
92 root.after(0, lambda: messagebox.showerror("Error", f"Generation failed:\n{e}"))
93 finally:
94 # Re-enable button and reset status
95 root.after(0, lambda: generate_button.config(state=tk.NORMAL, text="Generate"))
96 root.after(0, lambda: status_label.config(text="Ready."))
97
98def display_output(text):
99 """Insert generated text into the output box."""
100 output_text.delete("1.0", tk.END)
101 output_text.insert(tk.END, text)
102
103# ------------------- GUI Setup -------------------
104root = tk.Tk()
105root.title("GPT-5.1 0.4B Text Generator")
106root.geometry("800x700")
107root.resizable(True, True)
108
109# Style
110root.option_add("*Font", "Arial 11")
111
112# Status bar at the top
113status_label = tk.Label(root, text="Initializing...", bd=1, relief=tk.SUNKEN, anchor=tk.W)
114status_label.pack(fill=tk.X, padx=5, pady=2)
115
116# Main frame
117main_frame = ttk.Frame(root, padding="10")
118main_frame.pack(fill=tk.BOTH, expand=True)
119
120# Input label and text area
121ttk.Label(main_frame, text="Input Prompt:").grid(row=0, column=0, sticky=tk.W, pady=5)
122input_text = scrolledtext.ScrolledText(main_frame, height=8, wrap=tk.WORD)
123input_text.grid(row=1, column=0, columnspan=3, sticky=(tk.W, tk.E, tk.N, tk.S), pady=5)
124
125# Generation parameters
126param_frame = ttk.LabelFrame(main_frame, text="Generation Parameters", padding="10")
127param_frame.grid(row=2, column=0, columnspan=3, sticky=(tk.W, tk.E), pady=10)
128
129# Max new tokens
130ttk.Label(param_frame, text="Max New Tokens:").grid(row=0, column=0, sticky=tk.W, padx=5)
131max_new_tokens_var = tk.StringVar(value="100")
132ttk.Entry(param_frame, textvariable=max_new_tokens_var, width=8).grid(row=0, column=1, sticky=tk.W)
133
134# Temperature
135ttk.Label(param_frame, text="Temperature:").grid(row=0, column=2, sticky=tk.W, padx=5)
136temperature_var = tk.StringVar(value="0.7")
137ttk.Entry(param_frame, textvariable=temperature_var, width=8).grid(row=0, column=3, sticky=tk.W)
138
139# Top-p
140ttk.Label(param_frame, text="Top-p:").grid(row=1, column=0, sticky=tk.W, padx=5)
141top_p_var = tk.StringVar(value="0.9")
142ttk.Entry(param_frame, textvariable=top_p_var, width=8).grid(row=1, column=1, sticky=tk.W)
143
144# Repetition penalty
145ttk.Label(param_frame, text="Rep. Penalty:").grid(row=1, column=2, sticky=tk.W, padx=5)
146rep_penalty_var = tk.StringVar(value="1.1")
147ttk.Entry(param_frame, textvariable=rep_penalty_var, width=8).grid(row=1, column=3, sticky=tk.W)
148
149# Generate button
150generate_button = ttk.Button(main_frame, text="Generate", command=generate_text, state=tk.DISABLED)
151generate_button.grid(row=3, column=0, pady=10)
152
153# Output label and text area
154ttk.Label(main_frame, text="Generated Output:").grid(row=4, column=0, sticky=tk.W, pady=5)
155output_text = scrolledtext.ScrolledText(main_frame, height=15, wrap=tk.WORD)
156output_text.grid(row=5, column=0, columnspan=3, sticky=(tk.W, tk.E, tk.N, tk.S), pady=5)
157
158# Configure grid weights so text areas expand
159main_frame.columnconfigure(0, weight=1)
160main_frame.rowconfigure(1, weight=1)
161main_frame.rowconfigure(5, weight=1)
162
163# Start loading the model in the background after GUI is up
164root.after(100, lambda: threading.Thread(target=load_model, daemon=True).start())
165
166root.mainloop()