-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdata_utils.py
More file actions
144 lines (117 loc) · 4.94 KB
/
Copy pathdata_utils.py
File metadata and controls
144 lines (117 loc) · 4.94 KB
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
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
import os
from pathlib import Path
from torch.utils.data import Dataset, DataLoader, random_split
from PIL import Image
import pandas as pd
import albumentations as A
from albumentations.pytorch import ToTensorV2
import numpy as np
import torch
import argparse
class ClassificationDataset(Dataset):
"""
Custom Dataset for multi-label classification.
Assumes:
- Images are in <root>/train_images/
- Labels are in <root>/train.csv with columns ['ImageId', 'ClassId']
- ClassId: 0–4 (0 = no defect, 1–4 = defect types)
"""
def __init__(self, root, transform=None):
self.root = Path(root)
self.transform = transform
self.image_dir = self.root / "train_images"
self.image_paths = list(self.image_dir.glob("*.jpg"))
self.df = pd.read_csv(self.root / "train.csv")
self.df['ClassId'] = self.df['ClassId'].astype(int)
self.num_classes = 5
self.labels_dict = self._create_multilabel_dict()
def _create_multilabel_dict(self):
from collections import defaultdict
labels = defaultdict(lambda: [0] * self.num_classes)
for _, row in self.df.iterrows():
labels[row["ImageId"]][row["ClassId"]] = 1
return dict(labels)
def __len__(self):
return len(self.image_paths)
def __getitem__(self, index):
image_path = self.image_paths[index]
image = np.array(Image.open(image_path).convert("RGB"))
if self.transform:
image = self.transform(image=image)["image"]
label = torch.tensor(
self.labels_dict.get(image_path.name, [1, 0, 0, 0, 0]),
dtype=torch.float32
)
return image, label
class DefectBlackout(A.ImageOnlyTransform):
"""
Custom augmentation: Randomly black out patches to simulate defects.
"""
def __init__(self, always_apply=False, p=0.5):
super().__init__(always_apply, p)
def apply(self, img, **params):
h, w = img.shape[:2]
for _ in range(np.random.randint(1, 3)):
x1 = np.random.randint(0, w // 2)
y1 = np.random.randint(0, h // 2)
x2 = x1 + np.random.randint(w // 10, w // 4)
y2 = y1 + np.random.randint(h // 10, h // 4)
img[y1:y2, x1:x2] = 0
return img
def get_data_loaders(data_dir, batch_size=32, img_size=224, defect_blackout=True):
"""
Returns train and validation dataloaders.
Args:
data_dir (str): Root directory containing train.csv and train_images/
batch_size (int): Batch size
img_size (int): Height of crop (width is fixed at 1568)
defect_blackout (bool): If True, apply DefectBlackout augmentation
"""
transform_list = [
A.RandomCrop(height=img_size, width=1568),
A.HorizontalFlip(p=0.6),
A.VerticalFlip(p=0.6),
A.RandomBrightnessContrast(p=0.6),
]
if defect_blackout:
transform_list.append(DefectBlackout(p=0.5))
transform_list.extend([
A.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),
ToTensorV2()
])
transform = A.Compose(transform_list)
dataset = ClassificationDataset(root=data_dir, transform=transform)
train_size = int(0.9 * len(dataset))
val_size = len(dataset) - train_size
train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
train_dataset.dataset.transform = transform
val_dataset.dataset.transform = transform
num_workers = os.cpu_count()
train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)
val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)
return train_loader, val_loader
def main():
parser = argparse.ArgumentParser(description="DataLoader Preparation Script")
parser.add_argument('--data_dir', type=str, required=True,
help='Root directory containing train.csv and train_images/')
parser.add_argument('--batch_size', type=int, default=32,
help='Batch size for dataloaders')
parser.add_argument('--img_size', type=int, default=224,
help='Height for random crop (width fixed at 1568)')
parser.add_argument('--defect_blackout', action='store_true',
help='Apply DefectBlackout augmentation')
args = parser.parse_args()
train_loader, val_loader = get_data_loaders(
data_dir=args.data_dir,
batch_size=args.batch_size,
img_size=args.img_size,
defect_blackout=args.defect_blackout
)
print(f"\n Loaded data from {args.data_dir}")
print(f"Number of training samples: {len(train_loader.dataset)}")
print(f"Number of validation samples: {len(val_loader.dataset)}")
print(f"Batch size: {args.batch_size}")
print(f"Image crop size: {args.img_size} x 1568")
print(f"Defect blackout: {'enabled' if args.defect_blackout else 'disabled'}")
if __name__ == "__main__":
main()