CoolFace
Apppublic

nick-localhost/Sign-language-detection

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
data.py128 linesDownload Raw Back to src
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()