VJyzCELERY/ObjectClassificationPlayground
0
1from torch.utils.data import Subset,Dataset2import torch3import os4import numpy as np5import cv26 7def collate_fn(batch):8 imgs = [img for img, _ in batch]9 labels = torch.tensor([label for _, label in batch])10 return imgs, labels11 12 13class ImageDataset(Dataset):14 def __init__(self,root_path : str,img_size=(256,256)):15 classes = os.listdir(root_path)16 self.img_size = img_size17 self.classes = classes18 data = []19 for idx,class_name in enumerate(classes):20 class_path = os.path.join(root_path,class_name)21 files = os.listdir(class_path)22 for file in files:23 filepath = os.path.join(class_path,file)24 data.append({"image_path":filepath,"label":class_name,"id":idx})25 self.data = data26 27 def __len__(self):28 return len(self.data)29 30 def __getitem__(self,idx):31 curr = self.data[idx]32 label = curr['id']33 img_path = curr['image_path']34 img = cv2.imread(img_path)35 img = cv2.resize(img,(self.img_size))36 img = cv2.cvtColor(img,cv2.COLOR_BGR2RGB)37 img = img.astype(np.float32) / 255.038 return img,label39 40def simple_augment(img):41 if np.random.rand() > 0.5:42 img = cv2.flip(img, 1)43 44 angle = np.random.uniform(-15, 15)45 h, w = img.shape[:2]46 M = cv2.getRotationMatrix2D((w/2, h/2), angle, 1.0)47 img = cv2.warpAffine(img, M, (w, h), borderMode=cv2.BORDER_REFLECT)48 49 return img50 51 52class AugmentedSubset(Subset):53 def __init__(self, subset, augment_fn=None):54 super().__init__(subset.dataset, subset.indices)55 self.augment_fn = augment_fn56 57 def __getitem__(self, idx):58 img, label = super().__getitem__(idx)59 if self.augment_fn:60 img = self.augment_fn(img)61 return img, label62 