nick-localhost/Sign-language-detection
0
1import torch
2import numpy as np
3from torch.utils.data import DataLoader, Dataset
4import os
5from PIL import Image
6import albumentations as A
7import numpy as np
8from colorama import Fore
9from matplotlib import pyplot as plt
10from utils.boxes import rescale_bboxes, stacker
11from utils.setup import get_classes
12from utils.logger import get_logger
13from utils.rich_handlers import DataLoaderHandler
14import sys
15
16
17class DETRData(Dataset):
18 def __init__(self, path, train=True):
19 super().__init__()
20 self.path = path
21 self.labels_path = os.path.join(self.path, 'labels')
22 self.images_path = os.path.join(self.path, 'images')
23 self.label_files = os.listdir(self.labels_path)
24 self.labels = list(filter(lambda x: x.endswith('.txt'), self.label_files))
25 self.train = train
26
27 # Initialize logger
28 self.logger = get_logger("data_loader")
29 self.data_handler = DataLoaderHandler()
30
31 # Log dataset initialization
32 dataset_info = {
33 "Dataset Path": self.path,
34 "Mode": "Training" if train else "Testing",
35 "Total Samples": len(self.labels),
36 "Images Path": self.images_path,
37 "Labels Path": self.labels_path
38 }
39 self.data_handler.log_dataset_stats(dataset_info)
40
41 # Log transforms information
42 transform_list = [
43 "Resize to 500x500",
44 "Random Crop 224x224 (training only)",
45 "Final Resize to 224x224",
46 "Horizontal Flip p=0.5 (training only)",
47 "Color Jitter (training only)",
48 "Normalize (ImageNet stats)",
49 "Convert to Tensor"
50 ]
51 self.data_handler.log_transform_info(transform_list)
52
53 def safe_transform(self, image, bboxes, labels, max_attempts=50):
54 self.transform = A.Compose(
55 [
56 A.Resize(500,500),
57 *([A.RandomCrop(width=224, height=224, p=0.33)] if self.train else []), # Example random crop
58 A.Resize(224,224),
59 *([A.HorizontalFlip(p=0.5)] if self.train else []),
60 *([A.ColorJitter(brightness=0.5, contrast=0.5, saturation=0.5, hue=0.5, p=0.5)] if self.train else []),
61 A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
62 A.ToTensorV2()
63 ], bbox_params=A.BboxParams(format='yolo', label_fields=['class_labels'])
64 )
65
66 for attempt in range(max_attempts):
67 try:
68 transformed = self.transform(image=image, bboxes=bboxes, class_labels=labels)
69 # Check if we still have bboxes after transformation
70 if len(transformed['bboxes']) > 0:
71 return transformed
72 except:
73 continue
74
75 return {'image': image, 'bboxes': bboxes, 'class_labels': labels}
76
77 def __len__(self):
78 return len(self.labels)
79
80 def __getitem__(self, idx):
81 self.label_path = os.path.join(self.labels_path, self.labels[idx])
82 self.image_name = self.labels[idx].split('.')[0]
83 self.image_path = os.path.join(self.images_path, f'{self.image_name}.jpg')
84
85 img = Image.open(self.image_path)
86 with open(self.label_path, 'r') as f:
87 annotations = f.readlines()
88 class_labels = []
89 bounding_boxes = []
90 for annotation in annotations:
91 annotation = annotation.split('\n')[:-1][0].split(' ')
92 class_labels.append(annotation[0])
93 bounding_boxes.append(annotation[1:])
94 class_labels = np.array(class_labels).astype(int)
95 bounding_boxes = np.array(bounding_boxes).astype(float)
96
97 augmented = self.safe_transform(image=np.array(img), bboxes=bounding_boxes, labels=class_labels)
98 augmented_img_tensor = augmented['image']
99 augmented_bounding_boxes = np.array(augmented['bboxes'])
100 augmented_classes = augmented['class_labels']
101
102 labels = torch.tensor(augmented_classes, dtype=torch.long)
103 boxes = torch.tensor(augmented_bounding_boxes, dtype=torch.float32)
104 return augmented_img_tensor, {'labels': labels, 'boxes': boxes}
105
106if __name__ == '__main__':
107 dataset = DETRData('data/train', train=True)
108 dataloader = DataLoader(dataset, collate_fn=stacker, batch_size=4, drop_last=True)
109
110 X, y = next(iter(dataloader))
111 print(Fore.LIGHTCYAN_EX + str(y) + Fore.RESET)
112 CLASSES = get_classes()
113 fig, ax = plt.subplots(2,2)
114 axs = ax.flatten()
115 for idx, (img, annotations, ax) in enumerate(zip(X, y, axs)):
116 ax.imshow(img.permute(1,2,0))
117 box_classes = annotations['labels']
118 boxes = rescale_bboxes(annotations['boxes'], (224,224))
119 for box_class, bbox in zip(box_classes, boxes):
120 if box_class != 3:
121 xmin, ymin, xmax, ymax = bbox.detach().numpy()
122 print(xmin, ymin, xmax, ymax)
123 ax.add_patch(plt.Rectangle((xmin, ymin), xmax - xmin, ymax - ymin, fill=False, color=(0.000, 0.447, 0.741), linewidth=3))
124 text = f'{CLASSES[box_class]}'
125 ax.text(xmin, ymin, text, fontsize=15, bbox=dict(facecolor='yellow', alpha=0.5))
126
127 fig.tight_layout()
128 plt.show() 