Team Ai
Apppublic

binuser007/Toxic_comment_classification_using_Bert

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
data_loader.py106 linesDownload Raw Back to data
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)}")