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 "" in masked_text: detected_pii.append("full_name") if "" in masked_text: detected_pii.append("email") if "" in masked_text: detected_pii.append("phone_number") if "" in masked_text: detected_pii.append("dob") if "" in masked_text: detected_pii.append("aadhar") if "" in masked_text: detected_pii.append("credit/debit card") if "" in masked_text: detected_pii.append("cvv") if "" 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()