Shilpi Kumari
Add application file
a1ca7f2
Raw History Blame
2.39 kB
# inference.py - Simple inference script
import torch
from transformers import AutoTokenizer, AutoModelForSequenceClassification
import gradio as gr
class IncidentClassifier:
def __init__(self, model_name="Rajeshwartiwari/incident-classification-model"):
self.model_name = model_name
self.tokenizer = None
self.model = None
self.id2label = None
def load(self):
"""Load the model and tokenizer"""
print(f"Loading model: {self.model_name}")
self.tokenizer = AutoTokenizer.from_pretrained(self.model_name)
self.model = AutoModelForSequenceClassification.from_pretrained(self.model_name)
self.id2label = self.model.config.id2label
print(f"Model loaded with {len(self.id2label)} categories")
return self
def predict(self, text):
"""Make prediction on input text"""
if not text.strip():
return "Please enter text", {}
# Tokenize
inputs = self.tokenizer(
text,
return_tensors="pt",
truncation=True,
padding=True,
max_length=128
)
# Predict
with torch.no_grad():
outputs = self.model(**inputs)
probabilities = torch.nn.functional.softmax(outputs.logits, dim=-1)
# Get top prediction
predicted_id = outputs.logits.argmax().item()
predicted_label = self.id2label.get(predicted_id, "Unknown")
# Get all confidences
confidences = {}
for idx, prob in enumerate(probabilities[0]):
label = self.id2label.get(idx, f"Class_{idx}")
confidences[label] = float(prob) * 100
return predicted_label, confidences
# Quick test
if __name__ == "__main__":
classifier = IncidentClassifier().load()
test_cases = [
"Oracle database connection error ORA-12154",
"Email not syncing in Outlook",
"VPN keeps disconnecting every few minutes",
"Monitor screen flickering issues"
]
for test in test_cases:
label, confidences = classifier.predict(test)
top_3 = dict(sorted(confidences.items(), key=lambda x: x[1], reverse=True)[:3])
print(f"\n📝 Input: {test}")
print(f" 🎯 Predicted: {label}")
print(f" 📊 Top 3: {top_3}")