Download inference.py from Rajeshwartiwari/incident-classification-system: direct link, hf CLI and curl.
- Browser
- Download file 2.39 kB
-
https://huggingface.co/spaces/Rajeshwartiwari/incident-classification-system/resolve/a1ca7f27f2b4821eb40877f738fee8994c137478/inference.py
- Command line
-
hf download hf://spaces/Rajeshwartiwari/incident-classification-system@a1ca7f27f2b4821eb40877f738fee8994c137478/inference.py
-
curl -L -o inference.py https://huggingface.co/spaces/Rajeshwartiwari/incident-classification-system/resolve/a1ca7f27f2b4821eb40877f738fee8994c137478/inference.py
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}") |