Spaces:
Build error
Build error
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()
|