wenjun99 commited on
Commit
f8c3d9d
·
verified ·
1 Parent(s): de28637

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +18 -12
app.py CHANGED
@@ -1,21 +1,27 @@
1
- import torch
2
  from transformers import GPT2Tokenizer, GPT2LMHeadModel
3
 
4
- # Check if GPU is available, otherwise use CPU
5
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
6
-
7
- print(f"Device: {device}")
8
- # Load Fine-tuned Model and Tokenizer from Hugging Face Hub
9
- model = GPT2LMHeadModel.from_pretrained("wenjun99/gpt2-finetuned").to(device)
10
  tokenizer = GPT2Tokenizer.from_pretrained("wenjun99/gpt2-finetuned")
11
 
12
- # Inference Function
13
  def generate_response(query):
14
  input_text = f"Query: {query}\nTask:"
15
- inputs = tokenizer(input_text, return_tensors="pt").to(device)
16
  outputs = model.generate(**inputs, max_length=24, pad_token_id=tokenizer.eos_token_id)
17
  return tokenizer.decode(outputs[0], skip_special_tokens=True)
18
 
19
- # Test Example
20
- query = "Make a protein resistant to low-temperature environments?"
21
- print(generate_response(query))
 
 
 
 
 
 
 
 
 
 
 
1
+ import gradio as gr
2
  from transformers import GPT2Tokenizer, GPT2LMHeadModel
3
 
4
+ # Load Fine-tuned GPT-2 Model from Hugging Face
5
+ model = GPT2LMHeadModel.from_pretrained("wenjun99/gpt2-finetuned")
 
 
 
 
6
  tokenizer = GPT2Tokenizer.from_pretrained("wenjun99/gpt2-finetuned")
7
 
8
+ # Define Response Generation Function
9
  def generate_response(query):
10
  input_text = f"Query: {query}\nTask:"
11
+ inputs = tokenizer(input_text, return_tensors="pt").to(model.device)
12
  outputs = model.generate(**inputs, max_length=24, pad_token_id=tokenizer.eos_token_id)
13
  return tokenizer.decode(outputs[0], skip_special_tokens=True)
14
 
15
+ # Gradio UI
16
+ with gr.Blocks() as demo:
17
+ gr.Markdown("# 🤖 Fine-Tuned GPT-2 Chatbot")
18
+ gr.Markdown("Enter a query to see how the fine-tuned GPT-2 model responds.")
19
+
20
+ query_input = gr.Textbox(label="Enter Query")
21
+ generate_btn = gr.Button("Generate Response")
22
+ output_text = gr.Textbox(label="Generated Response")
23
+
24
+ generate_btn.click(generate_response, inputs=query_input, outputs=output_text)
25
+
26
+ # Launch Gradio App
27
+ demo.launch()