Demo / api.py
Nikpatil's picture
Upload 5 files
4af2eee verified
Raw History Blame Contribute Delete
1.65 kB
from flask import Blueprint, request, jsonify
from transformers import DebertaV2Tokenizer, DebertaV2ForSequenceClassification
import torch
from utils import mask_pii
api_bp = Blueprint("api", __name__)
# Repo of Hugging Face Model Hub where Model is Pushed
REPO_ID = "Nikpatil/Email_classifier"
MAX_LENGTH = 256
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()
id2label = {0: "Incident", 1: "Request", 2: "Problem", 3: "Change"}
@api_bp.route("/classify", methods=["POST"])
def classify_email():
data = request.get_json()
email_body = data.get("email_body", "")
if not email_body:
return jsonify({"Error": "Email body field is required"}), 400
masked_email, entities = mask_pii(email_body)
inputs = tokenizer(
masked_email,
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]
return jsonify({
"input_email_body": email_body,
"list_of_masked_entities": entities,
"masked_email": masked_email,
"category_of_the_email": predicted_class
}), 200