Nikpatil commited on
Commit
84bebe0
·
verified ·
1 Parent(s): a3a23b1

Upload 5 files

Browse files
Files changed (5) hide show
  1. api.py +54 -0
  2. app.py +8 -0
  3. model.py +265 -0
  4. requirements.txt +11 -0
  5. utils.py +288 -0
api.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from flask import Blueprint, request, jsonify
2
+ from transformers import DebertaV2Tokenizer, DebertaV2ForSequenceClassification
3
+ import torch
4
+ from utils import mask_pii
5
+
6
+ api_bp = Blueprint("api", __name__)
7
+
8
+ # Repo of Hugging Face Model Hub where Model is Pushed
9
+ REPO_ID = "Nikpatil/Email_classifier"
10
+ MAX_LENGTH = 256
11
+
12
+ tokenizer = DebertaV2Tokenizer.from_pretrained(REPO_ID)
13
+ model = DebertaV2ForSequenceClassification.from_pretrained(REPO_ID)
14
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
15
+ model.to(device)
16
+ model.eval()
17
+
18
+ id2label = {0: "Incident", 1: "Request", 2: "Problem", 3: "Change"}
19
+
20
+ @api_bp.route("/classify", methods=["POST"])
21
+ def classify_email():
22
+ data = request.get_json()
23
+ email_body = data.get("email_body", "")
24
+
25
+ if not email_body:
26
+ return jsonify({"Error": "Email body field is required"}), 400
27
+
28
+ masked_email, entities = mask_pii(email_body)
29
+
30
+ inputs = tokenizer(
31
+ masked_email,
32
+ add_special_tokens=True,
33
+ max_length=MAX_LENGTH,
34
+ padding='max_length',
35
+ truncation=True,
36
+ return_tensors='pt'
37
+ )
38
+ inputs = {k: v.to(device) for k, v in inputs.items()}
39
+
40
+ with torch.no_grad():
41
+ outputs = model(**inputs)
42
+
43
+ probs = torch.nn.functional.softmax(outputs.logits, dim=1)[0]
44
+ predicted_class_id = torch.argmax(probs).item()
45
+ predicted_class = id2label[predicted_class_id]
46
+
47
+ return jsonify({
48
+ "input_email_body": email_body,
49
+ "list_of_masked_entities": entities,
50
+ "masked_email": masked_email,
51
+ "category_of_the_email": predicted_class
52
+ }), 200
53
+
54
+
app.py ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ from flask import Flask
2
+ from api import api_bp
3
+
4
+ app = Flask(__name__)
5
+ app.register_blueprint(api_bp)
6
+
7
+ if __name__ == "__main__":
8
+ app.run(host="0.0.0.0", port=7860)
model.py ADDED
@@ -0,0 +1,265 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import pandas as pd
3
+ import torch
4
+ import torch.nn as nn
5
+ from sklearn.metrics import classification_report, confusion_matrix
6
+ from sklearn.model_selection import train_test_split
7
+ from transformers import (
8
+ DebertaV2Tokenizer,
9
+ DebertaV2ForSequenceClassification,
10
+ TrainingArguments,
11
+ Trainer,
12
+ EarlyStoppingCallback
13
+ )
14
+ import matplotlib.pyplot as plt
15
+ import seaborn as sns
16
+ from torch.utils.data import Dataset
17
+
18
+ # Set seed for reproducibility
19
+ SEED = 42
20
+ torch.manual_seed(SEED)
21
+ np.random.seed(SEED)
22
+
23
+ # ------------------------ Data Preprocessing ------------------------
24
+
25
+ class EmailClassification(Dataset):
26
+ def __init__(self, texts, labels, tokenizer, max_length):
27
+ """Initialize dataset with texts, labels and tokenizer settings."""
28
+ self.texts = texts
29
+ self.labels = labels
30
+ self.tokenizer = tokenizer
31
+ self.max_length = max_length
32
+
33
+ def __len__(self):
34
+ """Return the number of sample in the dataset."""
35
+ return len(self.texts)
36
+
37
+ def __getitem__(self, idx):
38
+ """Get a single item from the dataset."""
39
+ text = str(self.texts[idx])
40
+ label = self.labels[idx]
41
+
42
+ encoding = self.tokenizer(
43
+ text,
44
+ add_special_tokens=True,
45
+ max_length=self.max_length,
46
+ truncation=True,
47
+ return_attention_mask=True,
48
+ return_tensors='pt',
49
+ padding='max_length'
50
+ )
51
+
52
+ return {
53
+ 'input_ids': encoding['input_ids'].flatten(),
54
+ 'attention_mask': encoding['attention_mask'].flatten(),
55
+ 'labels': torch.tensor(label, dtype=torch.long)
56
+ }
57
+
58
+ def compute_class_weights(labels):
59
+ """Compute class weights inversely proportional to class frequencies."""
60
+ class_counts = np.bincount(labels)
61
+ total_samples = len(labels)
62
+ class_weights = total_samples / (len(class_counts) * class_counts)
63
+ return torch.tensor(class_weights, dtype=torch.float)
64
+
65
+
66
+ # ------------------------ Model Setup ------------------------
67
+
68
+ class WeightedTrainer(Trainer):
69
+ """Custom trainer that uses weighted loss function for imbalanced data."""
70
+ def __init__(self, *args, class_weights=None, **kwargs):
71
+ super().__init__(*args, **kwargs)
72
+ self.class_weights = class_weights.to(self.model.device)
73
+ self.loss_fn = nn.CrossEntropyLoss(
74
+ weight=self.class_weights,
75
+ label_smoothing=0.1
76
+ )
77
+
78
+ def compute_loss(self, model, inputs, return_outputs=False):
79
+ """Compute loss using the weighted loss function."""
80
+ labels = inputs.get("labels")
81
+ outputs = model(**inputs)
82
+ logits = outputs.get("logits")
83
+ loss = self.loss_fn(logits, labels)
84
+ return (loss, outputs) if return_outputs else loss
85
+
86
+ def load_model_and_tokenizer(model_name, label2id, id2label):
87
+ """Load the model and tokenizer."""
88
+ tokenizer = DebertaV2Tokenizer.from_pretrained(model_name)
89
+ model = DebertaV2ForSequenceClassification.from_pretrained(
90
+ model_name,
91
+ num_labels=len(label2id),
92
+ id2label=id2label,
93
+ label2id=label2id,
94
+ )
95
+ model.gradient_checkpointing_enable()
96
+ return model, tokenizer
97
+
98
+ # ------------------------ Metrics ------------------------
99
+
100
+ def compute_metrics(pred):
101
+ """Compute metrics for evaluation."""
102
+ labels = pred.label_ids
103
+ preds = pred.predictions.argmax(-1)
104
+
105
+ report = classification_report(
106
+ labels,
107
+ preds,
108
+ target_names=list(label2id.keys()),
109
+ output_dict=True
110
+ )
111
+
112
+ results = {
113
+ 'accuracy': report['accuracy'],
114
+ 'f1_macro': report['macro avg']['f1-score'],
115
+ 'f1_weighted': report['weighted avg']['f1-score'],
116
+ 'precision_weighted': report['weighted avg']['precision'],
117
+ 'recall_weighted': report['weighted avg']['recall']
118
+ }
119
+
120
+ for cls_name, cls_id in label2id.items():
121
+ results[f'f1_{cls_name}'] = report[cls_name]['f1-score']
122
+
123
+ return results
124
+
125
+
126
+ # ------------------------ Training and Evaluation ------------------------
127
+
128
+ def train_model(df, model_name, label2id, id2label, max_length, batch_size, epochs, seed):
129
+ """Train and evaluate model using train/val/test split."""
130
+ df['label_id'] = df['type'].map(label2id)
131
+
132
+ # Split the data
133
+ train_texts, temp_texts, train_labels, temp_labels = train_test_split(
134
+ df['email_processed'].values,
135
+ df['label_id'].values,
136
+ test_size=0.3,
137
+ stratify=df['label_id'].values,
138
+ random_state=seed
139
+ )
140
+ val_texts, test_texts, val_labels, test_labels = train_test_split(
141
+ temp_texts,
142
+ temp_labels,
143
+ test_size=0.5,
144
+ stratify=temp_labels,
145
+ random_state=seed
146
+ )
147
+
148
+ # Create datasets
149
+ tokenizer = DebertaV2Tokenizer.from_pretrained(model_name)
150
+ train_dataset = EmailClassification(train_texts, train_labels, tokenizer, max_length)
151
+ val_dataset = EmailClassification(val_texts, val_labels, tokenizer, max_length)
152
+ test_dataset = EmailClassification(test_texts, test_labels, tokenizer, max_length)
153
+
154
+ # Compute class weights
155
+ class_weights = compute_class_weights(train_labels)
156
+
157
+ # Load the model
158
+ model, tokenizer = load_model_and_tokenizer(model_name, label2id, id2label)
159
+
160
+ # Training arguments
161
+ training_args = TrainingArguments(
162
+ output_dir="./results",
163
+ num_train_epochs=epochs,
164
+ per_device_train_batch_size=batch_size,
165
+ per_device_eval_batch_size=batch_size,
166
+ learning_rate=5e-5,
167
+ lr_scheduler_type="cosine",
168
+ warmup_steps=120,
169
+ weight_decay=0.01,
170
+ logging_dir="./logs",
171
+ logging_steps=100,
172
+ eval_strategy="epoch",
173
+ save_strategy="epoch",
174
+ load_best_model_at_end=True,
175
+ metric_for_best_model="f1_weighted",
176
+ greater_is_better=True,
177
+ save_total_limit=1,
178
+ seed=seed,
179
+ fp16=True,
180
+ gradient_accumulation_steps=4
181
+ )
182
+
183
+ # Trainer with Weighted Loss
184
+ trainer = WeightedTrainer(
185
+ model=model,
186
+ args=training_args,
187
+ train_dataset=train_dataset,
188
+ eval_dataset=val_dataset,
189
+ compute_metrics=compute_metrics,
190
+ callbacks=[EarlyStoppingCallback(early_stopping_patience=2)],
191
+ class_weights=class_weights
192
+ )
193
+
194
+ # Train the model
195
+ trainer.train()
196
+
197
+ # Evaluate on validation and test set
198
+ val_results = trainer.evaluate()
199
+ test_results = trainer.predict(test_dataset)
200
+ test_metrics = compute_metrics(test_results)
201
+
202
+ # Print results
203
+ print("\nValidation Results:")
204
+ for metric_name, value in val_results.items():
205
+ print(f"{metric_name}: {value:.4f}")
206
+
207
+ print("\nTest Results:")
208
+ for metric_name, value in test_metrics.items():
209
+ print(f"{metric_name}: {value:.4f}")
210
+
211
+ # Get predictions for test set
212
+ test_preds = np.argmax(test_results.predictions, axis=1)
213
+
214
+ print("\nTest Set Classification Report:")
215
+ print(classification_report(test_labels, test_preds, target_names=list(label2id.keys())))
216
+
217
+ # Plot confusion matrix
218
+ plt.figure(figsize=(10, 8))
219
+ cm = confusion_matrix(test_labels, test_preds)
220
+ sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=list(label2id.keys()), yticklabels=list(label2id.keys()))
221
+ plt.xlabel('Predicted')
222
+ plt.ylabel('True')
223
+ plt.title('Confusion Matrix - Test Set')
224
+ plt.tight_layout()
225
+ plt.savefig('confusion_matrix_test_set.png')
226
+ plt.close()
227
+
228
+ model.save_pretrained('./model')
229
+ tokenizer.save_pretrained('./model')
230
+
231
+ return model, tokenizer
232
+
233
+
234
+ # ------------------------ Main ------------------------
235
+
236
+ if __name__ == "__main__":
237
+ label2id = {
238
+ 'Incident': 0,
239
+ 'Request': 1,
240
+ 'Problem': 2,
241
+ 'Change': 3
242
+ }
243
+ id2label = {v: k for k, v in label2id.items()}
244
+
245
+ # Model and Training Configuration
246
+ MODEL_NAME = "microsoft/mdeberta-v3-base"
247
+ MAX_LENGTH = 128
248
+ BATCH_SIZE = 4
249
+ EPOCHS = 5
250
+ SEED = 42
251
+
252
+ # Load Dataset
253
+ df = pd.read_csv("D:\OG Project\Data\combined_emails_with_natural_pii.csv")
254
+
255
+ # Train and evaluate the model
256
+ model, tokenizer = train_model(
257
+ df=df,
258
+ model_name=MODEL_NAME,
259
+ label2id=label2id,
260
+ id2label=id2label,
261
+ max_length=MAX_LENGTH,
262
+ batch_size=BATCH_SIZE,
263
+ epochs=EPOCHS,
264
+ seed=SEED
265
+ )
requirements.txt ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ numpy>=1.21.0
2
+ pandas>=1.3.0
3
+ matplotlib>=3.4.0
4
+ seaborn>=0.11.0
5
+ scikit-learn>=1.0.0
6
+ torch>=1.12.0
7
+ transformers>=4.26.0
8
+ sentencepiece>=0.1.95
9
+ protobuf>=3.20.0
10
+ spacy>=3.2.0
11
+
utils.py ADDED
@@ -0,0 +1,288 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import spacy
2
+ import re
3
+ import logging
4
+
5
+
6
+ logging.basicConfig(level=logging.INFO)
7
+ logger = logging.getLogger(__name__)
8
+
9
+ # Define regex patterns for specific PII fields
10
+ FULL_NAME_PATTERN = r'My name is ([A-Za-z\s\.\-]+)[\.,]|Name\s*:\s*([A-Za-z\s\.\-]+)[\.,]'
11
+ EMAIL_PATTERN = r'You can reach me at ([A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,})|Email\s*:\s*([A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,})'
12
+ PHONE_PATTERN = r'My [Cc]ontact number is\s*([+\d\s\-\(\)\.]+)|Phone\s*:\s*([+\d\s\-\(\)\.]+)|(\+?\d{1,4}?[-.\s]?\(?\d{1,3}?\)?[-.\s]?\d{1,4}[-.\s]?\d{1,4}[-.\s]?\d{1,9})'
13
+ DOB_PATTERN = r'[Dd]ate of [Bb]irth\s*:?\s*(\d{1,2}[-/\.]\d{1,2}[-/\.]\d{2,4}|\d{2,4}[-/\.]\d{1,2}[-/\.]\d{1,2})'
14
+ AADHAR_PATTERN = r'[Aa]adhar(?:\s*[Cc]ard)?\s*(?:[Nn]umber)?\s*:?\s*(\d{4}\s*\d{4}\s*\d{4}|\d{12})'
15
+ CREDIT_DEBIT_PATTERN = r'[Cc](?:redit|ard)\s*(?:[Nn]umber)?\s*:?\s*(\d{4}[-\s]?\d{4}[-\s]?\d{4}[-\s]?\d{4}|\d{16})'
16
+ CVV_PATTERN = r'[Cc][Vv][Vv]\s*(?:[Nn]umber)?\s*:?\s*(\d{3,4})'
17
+ EXPIRY_PATTERN = r'[Ee]xpir(?:y|ation)\s*[Dd]ate\s*:?\s*(\d{1,2}[-/\.]\d{2,4}|\d{2}[-/\.]\d{2})'
18
+
19
+ def preprocess_text(text):
20
+ """Clean and normalize text before processing. So that It impoves Model Consistency.
21
+ Args :
22
+ text: Email Text to Preprocess before Masking.
23
+ """
24
+ # Handle multiple consecutive newlines
25
+ text = re.sub(r'\n{2,}', '\n', text)
26
+ # Replace all remaining newlines with space
27
+ text = text.replace('\n', ' ')
28
+ # Remove excessive whitespaces So It can reduce Number of tokens
29
+ text = re.sub(r'\s+', ' ', text)
30
+ # Replace common unicode whitespace variants
31
+ text = text.replace('\xa0', ' ')
32
+ # Remove quotes that might be present in pasted emails
33
+ text = text.replace('"', '')
34
+
35
+ return text.strip()
36
+
37
+ def extract_entities(email_text):
38
+ """
39
+ Extract PII entities from text using regex patterns and spaCy
40
+
41
+ Args:
42
+ email_text (str): The email text to extract entities from
43
+
44
+ Returns:
45
+ list: List of dictionaries containing entity information
46
+ """
47
+ if not email_text or len(email_text.strip()) == 0:
48
+ return []
49
+
50
+ # Load spaCy model
51
+ try:
52
+ nlp = spacy.load("en_core_web_sm")
53
+ logger.info("Successfully loaded Spacy model")
54
+ except OSError as e:
55
+ logger.error(f"Error loading Spacy model: {e}")
56
+ raise
57
+
58
+ preprocessed_text = preprocess_text(email_text)
59
+ entities = []
60
+
61
+ #Extract full names and also Postitons of Entity
62
+ name_matches = re.finditer(FULL_NAME_PATTERN, preprocessed_text)
63
+ for match in name_matches:
64
+ name = next((g for g in match.groups() if g), "")
65
+ if name:
66
+ start_idx = email_text.find(name)
67
+ if start_idx != -1:
68
+ entities.append({
69
+ "start": start_idx,
70
+ "end": start_idx + len(name),
71
+ "text": name,
72
+ "type": "full_name"
73
+ })
74
+
75
+ #Extract email addresses and also Postitons of Entity
76
+ email_matches = re.finditer(EMAIL_PATTERN, preprocessed_text)
77
+ for match in email_matches:
78
+ email = next((g for g in match.groups() if g), "")
79
+ if email:
80
+ start_idx = email_text.find(email)
81
+ if start_idx != -1:
82
+ entities.append({
83
+ "start": start_idx,
84
+ "end": start_idx + len(email),
85
+ "text": email,
86
+ "type": "email"
87
+ })
88
+
89
+ #Extract phone numbers and also Postitons of Entity
90
+ phone_matches = re.finditer(PHONE_PATTERN, preprocessed_text)
91
+ for match in phone_matches:
92
+ phone = next((g for g in match.groups() if g), "")
93
+ if phone:
94
+ start_idx = email_text.find(phone)
95
+ if start_idx != -1:
96
+ entities.append({
97
+ "start": start_idx,
98
+ "end": start_idx + len(phone),
99
+ "text": phone,
100
+ "type": "phone_number"
101
+ })
102
+
103
+ #Extract date of birth and also Postitons of Entity
104
+ dob_matches = re.finditer(DOB_PATTERN, preprocessed_text)
105
+ for match in dob_matches:
106
+ dob = match.group(1)
107
+ if dob:
108
+ start_idx = email_text.find(dob)
109
+ if start_idx != -1:
110
+ entities.append({
111
+ "start": start_idx,
112
+ "end": start_idx + len(dob),
113
+ "text": dob,
114
+ "type": "dob"
115
+ })
116
+
117
+ #Extract Aadhar numbers and also Postitons of Entity
118
+ aadhar_matches = re.finditer(AADHAR_PATTERN, preprocessed_text)
119
+ for match in aadhar_matches:
120
+ aadhar = match.group(1)
121
+ if aadhar:
122
+ start_idx = email_text.find(aadhar)
123
+ if start_idx != -1:
124
+ entities.append({
125
+ "start": start_idx,
126
+ "end": start_idx + len(aadhar),
127
+ "text": aadhar,
128
+ "type": "aadhar_num"
129
+ })
130
+
131
+ #Extract credit/debit card numbers and also Postitons of Entity
132
+ card_matches = re.finditer(CREDIT_DEBIT_PATTERN, preprocessed_text)
133
+ for match in card_matches:
134
+ card = match.group(1)
135
+ if card:
136
+ start_idx = email_text.find(card)
137
+ if start_idx != -1:
138
+ entities.append({
139
+ "start": start_idx,
140
+ "end": start_idx + len(card),
141
+ "text": card,
142
+ "type": "credit_debit_no"
143
+ })
144
+
145
+ #Extract CVV numbers and also Postitons of Entity
146
+ cvv_matches = re.finditer(CVV_PATTERN, preprocessed_text)
147
+ for match in cvv_matches:
148
+ cvv = match.group(1)
149
+ if cvv:
150
+ start_idx = email_text.find(cvv)
151
+ if start_idx != -1:
152
+ entities.append({
153
+ "start": start_idx,
154
+ "end": start_idx + len(cvv),
155
+ "text": cvv,
156
+ "type": "cvv_no"
157
+ })
158
+
159
+ #Extract card expiry dates and also Postitons of Entity
160
+ expiry_matches = re.finditer(EXPIRY_PATTERN, preprocessed_text)
161
+ for match in expiry_matches:
162
+ expiry = match.group(1)
163
+ if expiry:
164
+ start_idx = email_text.find(expiry)
165
+ if start_idx != -1:
166
+ entities.append({
167
+ "start": start_idx,
168
+ "end": start_idx + len(expiry),
169
+ "text": expiry,
170
+ "type": "expiry_no"
171
+ })
172
+
173
+ # Use spaCy for additional PII detection
174
+ doc = nlp(preprocessed_text)
175
+
176
+ # Used spaCy's more advanced NER capabilities
177
+ # Created a mapping from character indices in preprocessed_text to original email_text
178
+ char_mapping = {}
179
+ preprocessed_idx = 0
180
+ original_idx = 0
181
+
182
+ while preprocessed_idx < len(preprocessed_text) and original_idx < len(email_text):
183
+ if preprocessed_text[preprocessed_idx] == email_text[original_idx]:
184
+ char_mapping[preprocessed_idx] = original_idx
185
+ preprocessed_idx += 1
186
+ original_idx += 1
187
+ else:
188
+ # Skip characters in original text that were removed during preprocessing
189
+ original_idx += 1
190
+
191
+ # Process spaCy entities with proper position mapping
192
+ for ent in doc.ents:
193
+ entity_type = None
194
+
195
+ # Map spaCy entity types to our entity types
196
+ if ent.label_ == "PERSON":
197
+ entity_type = "full_name"
198
+ elif ent.label_ == "DATE":
199
+ # Check if it looks like a date of birth or expiry date
200
+ if re.search(DOB_PATTERN, ent.text):
201
+ entity_type = "dob"
202
+ elif re.search(EXPIRY_PATTERN, ent.text):
203
+ entity_type = "expiry_no"
204
+ elif ent.label_ == "CARDINAL" or ent.label_ == "MONEY":
205
+ # Check if it looks like a credit card, aadhar, or other sensitive numbers
206
+ if re.search(CREDIT_DEBIT_PATTERN, ent.text):
207
+ entity_type = "credit_debit_no"
208
+ elif re.search(AADHAR_PATTERN, ent.text):
209
+ entity_type = "aadhar_num"
210
+ elif re.search(CVV_PATTERN, ent.text):
211
+ entity_type = "cvv_no"
212
+ elif ent.label_ == "ORG" and "@" in ent.text:
213
+ # Sometimes spaCy identifies emails as organizations
214
+ entity_type = "email"
215
+
216
+ if entity_type:
217
+ # Check if this entity overlaps with any existing entity
218
+ already_captured = False
219
+ ent_start_idx = char_mapping.get(ent.start_char, -1)
220
+
221
+ if ent_start_idx != -1:
222
+ ent_end_idx = char_mapping.get(min(ent.end_char, len(preprocessed_text)-1),
223
+ ent_start_idx + len(ent.text))
224
+
225
+ # Check for overlap with existing entities
226
+ for entity in entities:
227
+ if (entity["start"] <= ent_end_idx and
228
+ entity["end"] >= ent_start_idx):
229
+ already_captured = True
230
+ break
231
+
232
+ if not already_captured:
233
+ # Get the actual text from the original email
234
+ original_text = email_text[ent_start_idx:ent_end_idx]
235
+
236
+ entities.append({
237
+ "start": ent_start_idx,
238
+ "end": ent_end_idx,
239
+ "text": original_text,
240
+ "type": entity_type
241
+ })
242
+
243
+ # Sort entities by start position
244
+ entities.sort(key=lambda x: x["start"])
245
+
246
+ return entities
247
+
248
+ def mask_pii(email_text):
249
+ """
250
+ Mask PII in email text and return both masked text and entity list.
251
+
252
+ Args:
253
+ email_text (str): Original email text
254
+
255
+ Returns:
256
+ tuple: (masked_text, list_of_masked_entities)
257
+ - masked_text: Text with PII replaced by entity type tags
258
+ - list_of_masked_entities: List of dictionaries with entity information
259
+ """
260
+ if not email_text or len(email_text.strip()) == 0:
261
+ return "", []
262
+
263
+ # Extract all entities from text
264
+ entities = extract_entities(email_text)
265
+
266
+ # Create a copy of the original text
267
+ masked_text = email_text
268
+
269
+ # Process entities in reverse order to avoid index shifting when replacing
270
+ for entity in sorted(entities, key=lambda x: x["start"], reverse=True):
271
+ start, end = entity["start"], entity["end"]
272
+ entity_type = entity["type"]
273
+ original_text = entity["text"]
274
+
275
+ # Replace the entity with a tag
276
+ masked_text = masked_text[:start] + f"<{entity_type}>" + masked_text[end:]
277
+
278
+ # Format entities for output
279
+ formatted_entities = []
280
+ for entity in entities:
281
+ formatted_entities.append({
282
+ "position": [entity["start"], entity["end"]],
283
+ "classification": entity["type"],
284
+ "entity": entity["text"]
285
+ })
286
+
287
+ return masked_text, formatted_entities
288
+