File size: 3,608 Bytes
69b17de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
"""
ML Pipeline Data Loader for ScamDetect AI
Handles the 50,000+ research-level text, URL, image, and video datasets.
Provides PyTorch Dataset and DataLoader abstractions.
"""
import os
import pandas as pd
try:
    import torch
    from torch.utils.data import Dataset, DataLoader
    from transformers import AutoTokenizer
    TORCH_AVAILABLE = True
except ImportError:
    TORCH_AVAILABLE = False
    print("Warning: PyTorch not installed. ML DataLoader will operate in pandas-only mode.")

DATA_DIR = os.path.join(os.path.dirname(__file__), "data")

class TextScamDataset(Dataset if TORCH_AVAILABLE else object):
    def __init__(self, csv_file=None, tokenizer_name="bert-base-multilingual-cased", max_length=128):
        self.csv_file = csv_file or os.path.join(DATA_DIR, "text_dataset.csv")
        self.data = pd.read_csv(self.csv_file)
        
        # Mapping labels
        self.label_map = {"safe": 0, "scam": 1}
        self.max_length = max_length
        
        if TORCH_AVAILABLE:
            self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
        
    def __len__(self):
        return len(self.data)
        
    def __getitem__(self, idx):
        row = self.data.iloc[idx]
        text = str(row['text'])
        label = self.label_map.get(row['label'], 1)
        
        if not TORCH_AVAILABLE:
            return {"text": text, "label": label, "category": row['category']}
            
        encoding = self.tokenizer(
            text,
            add_special_tokens=True,
            max_length=self.max_length,
            return_token_type_ids=False,
            padding='max_length',
            truncation=True,
            return_attention_mask=True,
            return_tensors='pt',
        )
        
        return {
            'text': text,
            'input_ids': encoding['input_ids'].flatten(),
            'attention_mask': encoding['attention_mask'].flatten(),
            'targets': torch.tensor(label, dtype=torch.long),
            'category': row['category']
        }

class UrlPhishingDataset(Dataset if TORCH_AVAILABLE else object):
    def __init__(self, csv_file=None):
        self.csv_file = csv_file or os.path.join(DATA_DIR, "url_dataset.csv")
        self.data = pd.read_csv(self.csv_file)
        self.label_map = {"safe": 0, "scam": 1}
        
    def __len__(self):
        return len(self.data)
        
    def __getitem__(self, idx):
        row = self.data.iloc[idx]
        url = str(row['url'])
        label = self.label_map.get(row['label'], 1)
        
        item = {"url": url, "label": label}
        if TORCH_AVAILABLE:
            item["targets"] = torch.tensor(label, dtype=torch.long)
        return item

def get_data_loaders(batch_size=32):
    """Returns PyTorch DataLoaders for train/test splits."""
    if not TORCH_AVAILABLE:
        raise ImportError("PyTorch required for DataLoader generation")
        
    text_ds = TextScamDataset()
    url_ds = UrlPhishingDataset()
    
    # In a real scenario, split train/test here using torch.utils.data.random_split
    text_loader = DataLoader(text_ds, batch_size=batch_size, shuffle=True)
    url_loader = DataLoader(url_ds, batch_size=batch_size, shuffle=True)
    
    return {"text": text_loader, "url": url_loader}

if __name__ == "__main__":
    print("Loading datasets...")
    text_ds = TextScamDataset()
    url_ds = UrlPhishingDataset()
    print(f"Loaded Text Dataset: {len(text_ds)} samples")
    print(f"Loaded URL Dataset: {len(url_ds)} samples")
    
    if len(text_ds) > 0:
        print("\nSample Text Record:")
        print(text_ds[0])