Atharva_Chimera / Version_2 /core /data_loader.py
hch-dev
Reorganized and renamed Version folders
e6bbf5e
Raw
History Blame Contribute Delete
2.52 kB
import torch
import pandas as pd
from torch.utils.data import DataLoader, Dataset
from PIL import Image
from core.augmentations import get_train_transforms, get_eval_transforms
class PhishingImageDataset(Dataset):
def __init__(self, dataframe, transform=None, image_col='image_path', label_col='label'):
self.dataframe = dataframe
self.transform = transform
self.image_col = image_col
self.label_col = label_col
def __len__(self):
return len(self.dataframe)
def __getitem__(self, idx):
row = self.dataframe.iloc[idx]
img_path = row[self.image_col]
label = row[self.label_col]
try:
image = Image.open(img_path).convert("RGB")
except Exception as e:
# Fallback to a blank image if file is missing/corrupted
print(f"Error loading {img_path}: {e}")
image = Image.new('RGB', (224, 224), (0, 0, 0))
if self.transform:
image = self.transform(image)
return image, label
def prepare_dataloaders(legit_csv_path, phishing_csv_path, batch_size=32, image_column_name='image_path'):
print(f"Sampling 5,000 rows from {legit_csv_path} and {phishing_csv_path}...")
df_legit = pd.read_csv(legit_csv_path).sample(n=5000, random_state=42)
df_legit['label'] = 0 # 0: Legit
df_phish = pd.read_csv(phishing_csv_path).sample(n=5000, random_state=42)
df_phish['label'] = 1 # 1: Phishing
df_all = pd.concat([df_legit, df_phish], ignore_index=True)
df_all = df_all.sample(frac=1, random_state=42).reset_index(drop=True)
# 80/10/10 Split -> 8000 Train, 1000 Val, 1000 Test
train_df = df_all.iloc[:8000]
val_df = df_all.iloc[8000:9000]
test_df = df_all.iloc[9000:]
train_ds = PhishingImageDataset(train_df, transform=get_train_transforms(), image_col=image_column_name)
val_ds = PhishingImageDataset(val_df, transform=get_eval_transforms(), image_col=image_column_name)
test_ds = PhishingImageDataset(test_df, transform=get_eval_transforms(), image_col=image_column_name)
train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=4)
val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=4)
test_loader = DataLoader(test_ds, batch_size=batch_size, shuffle=False, num_workers=4)
return train_loader, val_loader, test_loader, 2