Spaces:
Paused
Paused
Download app.py from Nikpatil/classifier: direct link, hf CLI and curl.
- Browser
- Download file 2.52 kB
-
https://huggingface.co/spaces/Nikpatil/classifier/resolve/main/app.py
- Command line
-
hf download hf://spaces/Nikpatil/classifier/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Nikpatil/classifier/resolve/main/app.py
2.52 kB
| import torch | |
| from transformers import DebertaV2Tokenizer, DebertaV2ForSequenceClassification | |
| import gradio as gr | |
| from preprocess import mask_pii | |
| REPO_ID = "Nikpatil/Email_classifier" | |
| MAX_LENGTH = 256 | |
| # Define your label mapping | |
| id2label = {0: "Incident", 1: "Request", 2: "Problem", 3: "Change"} | |
| labels = list(id2label.values()) | |
| # Load model & tokenizer | |
| tokenizer = DebertaV2Tokenizer.from_pretrained(REPO_ID) | |
| model = DebertaV2ForSequenceClassification.from_pretrained(REPO_ID) | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model.to(device) | |
| model.eval() | |
| # Your inference function (adapted from your own) | |
| def inference(text, model, tokenizer, max_length, id2label): | |
| masked_text = mask_pii(text) | |
| inputs = tokenizer( | |
| masked_text, | |
| add_special_tokens=True, | |
| max_length=MAX_LENGTH, | |
| padding='max_length', | |
| truncation=True, | |
| return_tensors='pt' | |
| ) | |
| inputs = {k: v.to(device) for k, v in inputs.items()} | |
| with torch.no_grad(): | |
| outputs = model(**inputs) | |
| probs = torch.nn.functional.softmax(outputs.logits, dim=1)[0] | |
| predicted_class_id = torch.argmax(probs).item() | |
| predicted_class = id2label[predicted_class_id] | |
| # Step 4: Extract PII fields (back from masked text) | |
| detected_pii = [] | |
| if "<full_name>" in masked_text: | |
| detected_pii.append("full_name") | |
| if "<email>" in masked_text: | |
| detected_pii.append("email") | |
| if "<phone_number>" in masked_text: | |
| detected_pii.append("phone_number") | |
| if "<dob>" in masked_text: | |
| detected_pii.append("dob") | |
| if "<aadhar_num>" in masked_text: | |
| detected_pii.append("aadhar") | |
| if "<credit_debit_no>" in masked_text: | |
| detected_pii.append("credit/debit card") | |
| if "<cvv_no>" in masked_text: | |
| detected_pii.append("cvv") | |
| if "<expiry_no>" in masked_text: | |
| detected_pii.append("expiry date") | |
| return { | |
| "original_email": text, | |
| "masked_email": masked_text, | |
| "predicted_email_class": predicted_class, | |
| "prediction_confidence": float(probs[predicted_class_id]), | |
| "detected_pii": detected_pii | |
| } | |
| # Gradio interface | |
| demo = gr.Interface( | |
| fn=lambda text: inference(text, model, tokenizer, MAX_LENGTH, id2label), | |
| inputs=gr.Textbox(lines=8, placeholder="Paste your email here..."), | |
| outputs="json", | |
| title="Email Classifier", | |
| description="Classifies emails into Incident, Request, Problem, or Change." | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch() | |