File size: 1,776 Bytes
2fc306d
 
b7b2be5
 
 
 
2fc306d
 
 
 
b7b2be5
 
 
 
 
 
 
 
f7745b3
b7b2be5
f7745b3
b7b2be5
 
 
 
 
 
 
 
 
 
 
 
 
90c9a78
 
 
 
 
b7b2be5
 
 
 
2fc306d
b7b2be5
 
f7745b3
b7b2be5
f7745b3
 
b7b2be5
 
 
2fc306d
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
import os
import shutil
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
import gradio as gr

# Check if model is extracted; if not, extract it
if not os.path.exists("best_model"):
    shutil.unpack_archive("best_model.zip", "best_model")

# Load the saved model and tokenizer
model = AutoModelForSequenceClassification.from_pretrained("best_model")
tokenizer = AutoTokenizer.from_pretrained("best_model")

# Ensure the model is in evaluation mode
model.eval()

# Define the prediction function
def predict(Text):
    # Tokenize the input text
    inputs = tokenizer(Text, return_tensors="pt", max_length=512, truncation=True, padding=True)
    
    # Perform inference
    with torch.no_grad():
        logits = model(**inputs).logits
    
    # Get predicted label and confidence scores
    probs = torch.nn.functional.softmax(logits, dim=1)
    _, predicted_label = torch.max(logits, dim=1)
    
    # Map the predicted label to a human-readable class name
    class_names = ["Class a", "Class b", "Class c", "Class d", "Class e"]  
    predicted_class = class_names[predicted_label.item()]
    
    # Convert confidence scores to percentage with 2 decimal places
    probs_percentage = [f"{p * 100:.2f}%" for p in probs.tolist()[0]]
    
    # Return the predicted class and formatted confidence scores
    return predicted_class, str(probs_percentage)

# Create the Gradio interface
iface = gr.Interface(
    fn=predict,
    inputs=gr.Textbox(lines=2, placeholder="Enter your text"),
    outputs=[
        gr.Textbox(label="Predicted Class"),
        gr.Textbox(label="Confidence")
    ],
    title="Chronological Classification",
    description="Classify poem into predefined categories."
)

# Launch the interface
iface.launch()