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()