classifier / app.py
Nikpatil's picture
Update app.py
4921c6b verified
Raw History Blame Contribute Delete
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()