Spaces:
Paused
Paused
File size: 2,522 Bytes
2494037 4921c6b 2494037 39cc5cb 5383e20 2494037 5383e20 2494037 39cc5cb 2494037 5383e20 2494037 5383e20 2494037 | 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 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 | 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()
|