binuser007/Toxic_comment_classification_using_Bert
0
1import pandas as pd
2import torch
3from torch.utils.data import Dataset, DataLoader
4from transformers import BertTokenizer
5from typing import Dict, List, Tuple
6import numpy as np
7import os
8
9class ToxicCommentDataset(Dataset):
10 def __init__(self, texts: List[str], labels: np.ndarray, tokenizer: BertTokenizer, max_length: int = 128):
11 # Convert texts to list if it's a pandas Series
12 self.texts = texts.tolist() if isinstance(texts, pd.Series) else texts
13 self.labels = labels
14 self.tokenizer = tokenizer
15 self.max_length = max_length
16
17 def __len__(self):
18 return len(self.texts)
19
20 def __getitem__(self, idx) -> Dict[str, torch.Tensor]:
21 text = str(self.texts[idx])
22
23 # Handle unusual line terminators
24 text = text.replace('\u2028', ' ').replace('\u2029', ' ') # Remove line/paragraph separators
25 text = ' '.join(text.splitlines()) # Normalize all newlines
26
27 label = self.labels[idx]
28
29 encoding = self.tokenizer(
30 text,
31 add_special_tokens=True,
32 max_length=self.max_length,
33 padding='max_length',
34 truncation=True,
35 return_tensors='pt'
36 )
37
38 return {
39 'input_ids': encoding['input_ids'].flatten(),
40 'attention_mask': encoding['attention_mask'].flatten(),
41 'labels': torch.FloatTensor(label)
42 }
43
44def load_toxic_data(data_path: str) -> Tuple[List[str], np.ndarray]:
45 """Load and prepare the toxic comment dataset"""
46 try:
47 # Use encoding='utf-8-sig' to handle BOM if present
48 df = pd.read_csv(data_path, encoding='utf-8-sig', on_bad_lines='skip')
49
50 # List of toxicity categories
51 toxic_categories = ['toxic', 'severe_toxic', 'obscene', 'threat', 'insult', 'identity_hate']
52
53 # Convert text column to list and labels to numpy array
54 texts = df['comment_text'].tolist()
55 labels = df[toxic_categories].values
56
57 return texts, labels
58 except Exception as e:
59 raise RuntimeError(f"Error loading data from {data_path}: {str(e)}")
60
61def create_data_loaders(
62 texts: List[str],
63 labels: np.ndarray,
64 tokenizer: BertTokenizer,
65 train_ratio: float = 0.8,
66 batch_size: int = 32,
67 num_workers: int = 4 # Adjusted for Windows
68) -> Tuple[DataLoader, DataLoader]:
69 """Create train and validation data loaders"""
70 try:
71 # Calculate split index
72 dataset_size = len(texts)
73 train_size = int(dataset_size * train_ratio)
74
75 # Split data
76 train_texts = texts[:train_size]
77 train_labels = labels[:train_size]
78 val_texts = texts[train_size:]
79 val_labels = labels[train_size:]
80
81 # Create datasets
82 train_dataset = ToxicCommentDataset(train_texts, train_labels, tokenizer)
83 val_dataset = ToxicCommentDataset(val_texts, val_labels, tokenizer)
84
85 # Create data loaders with Windows-optimized settings
86 train_loader = DataLoader(
87 train_dataset,
88 batch_size=batch_size,
89 shuffle=True,
90 num_workers=num_workers,
91 pin_memory=True, # Helps with CUDA performance
92 persistent_workers=True # Keeps workers alive between epochs
93 )
94
95 val_loader = DataLoader(
96 val_dataset,
97 batch_size=batch_size,
98 shuffle=False,
99 num_workers=num_workers,
100 pin_memory=True,
101 persistent_workers=True
102 )
103
104 return train_loader, val_loader
105 except Exception as e:
106 raise RuntimeError(f"Error creating data loaders: {str(e)}")