from PIL import Image from typing import Any, Callable, Dict, List, Optional, Tuple from torchvision.datasets.folder import is_image_file, find_classes, pil_loader, VisionDataset import os import numpy as np import torch from collections import defaultdict def make_dataset(directory: str, class_to_idx: Optional[Dict[str, int]], frames_per_clip: int) -> List[Tuple[str, int]]: directory = os.path.expanduser(directory) instances = [] class_n_samples = defaultdict(int) for target_class in sorted(class_to_idx.keys()): class_index = class_to_idx[target_class] target_dir = os.path.join(directory, target_class) if not os.path.isdir(target_dir): continue video_folders = [item.name for item in os.scandir(target_dir) if item.is_dir()] video_folders.sort() for index, folder in enumerate(video_folders): if index % 1000 == 0: print(f'Processing video at index:{index} out of {len(video_folders)}') folder = os.path.join(target_dir, folder) imagefiles = [item for item in os.listdir(folder) if is_image_file(item)] num_images = len(imagefiles) # TODO: make sure imagefiles follow a given format. E.g. `0.jpg` to `{num_frames-1}.jpg` sample = folder, num_images, class_index if num_images >= frames_per_clip: instances.append(sample) class_n_samples[target_class] += 1 return instances, class_n_samples def _loader(folder_containing_frames: str, num_frames: int, frames_per_clip: int) -> List[Image.Image]: index = np.random.choice(num_frames - frames_per_clip + 1) imagefiles = [os.path.join(folder_containing_frames, f'{index + i}.jpg') for i in range(frames_per_clip)] return [pil_loader(imagefile) for imagefile in imagefiles] class DatasetFolder(VisionDataset): """A generic data loader. . ├── dataset1 │ ├── video_1 │ │ ├── 0.jpg │ │ ├── 1.jpg │ │ ├── 2.jpg │ │ ├── ... │ │ └── n1.jpg │ └── video_2 ├── dataset2 │ ├── video_1 │ │ ├── 0.jpg │ │ ├── 1.jpg │ │ ├── 2.jpg │ │ ├── ... │ │ └── n2.jpg │ └── video_2 └── dataset3 ├── video_1 │ ├── 0.jpg │ ├── 1.jpg │ ├── 2.jpg │ ├── ... │ └── n3.jpg └── video_2 """ def __init__( self, root: str, balanced_dataset: Optional[bool], loader: Callable[[str], Any] = _loader, frames_per_clip: int = 0, transform_train: Optional[Callable] = None, transform_val: Optional[Callable] = None, target_transform: Optional[Callable] = None, excluded_folders: List[str]= []) -> None: super().__init__(root, transform=None, target_transform=target_transform) assert frames_per_clip > 0 self.frames_per_clip = frames_per_clip self.transform_train, self.transform_val = transform_train, transform_val self.balanced_dataset = balanced_dataset self.val_indices = [] classes, class_to_idx = find_classes(self.root) for excluded_class in excluded_folders: if excluded_class in class_to_idx: del class_to_idx[excluded_class] classes.remove(excluded_class) assert len(classes) == len(class_to_idx) self.classes = classes self.class_to_idx = class_to_idx samples, class_n_samples = make_dataset(self.root, class_to_idx, frames_per_clip) self.samples = samples self.class_n_samples = class_n_samples self.loader = loader self.targets = [s[2] for s in samples] self.val_indices = None def index_to_transform(self, index: int): return self.transform_val if index in self.val_indices else self.transform_train def __getitem__(self, index: int) -> Tuple[Any, Any]: """ Output: a tuple (sample, target) where - sample is a tensor of size (frames_per_clip, 3, 224, 224) - `target` is class_index of the target class. """ if self.balanced_dataset: if index in self.val_indices: random_target = np.random.choice(list(self.target_to_indices_val.keys())) index = np.random.choice(self.target_to_indices_val[random_target]) else: random_target = np.random.choice(list(self.target_to_indices_train.keys())) index = np.random.choice(self.target_to_indices_train[random_target]) path, num_frames, target = self.samples[index] transform = self.index_to_transform(index) list_of_consecutive_frames = self.loader(path, num_frames, self.frames_per_clip) if transform is not None: sample = transform(list_of_consecutive_frames) sample = torch.stack(sample, dim=0) if self.target_transform is not None: target = self.target_transform(target) return sample, target def set_target_to_indices_dict(self, target_transform, samples): ''' samples: List of tuples where each tuple is (folder, num_images, class_index) target_transform: maps class_index to targets required for training ''' assert self.val_indices is not None target_to_indices_train, target_to_indices_val = defaultdict(list), defaultdict(list) for index, (_, _, class_index) in enumerate(samples): key_name = target_transform(class_index) if index in self.val_indices: target_to_indices_val[key_name].append(index) else: target_to_indices_train[key_name].append(index) self.target_to_indices_val = target_to_indices_val self.target_to_indices_train = target_to_indices_train @property def pos_weight(self): n_pos, n_neg = 0, 0 for class_name, n_samples in self.class_n_samples.items(): class_idx = self.class_to_idx[class_name] if self.target_transform(class_idx) == 1: n_pos += n_samples else: n_neg += n_samples assert n_pos > 0 and n_neg > 0 return n_neg/n_pos def __len__(self) -> int: return len(self.samples) def get_target_transform(class_to_idx: Optional[Dict[str, int]], *positive_labels): positive_target = [class_to_idx[label] for label in positive_labels] def target_transform(target): return int(target in positive_target) return target_transform if __name__ == '__main__': from torch.utils.data import DataLoader from torch.utils.data.sampler import SubsetRandomSampler from batch_image_transforms import batch_transform_train, batch_transform_val DATASET_ROOT = '../mock_videoframes_dataset' print('Testing the dataset definition now') positive_labels = ['dataset1', 'dataset2'] ds = DatasetFolder(DATASET_ROOT, loader=_loader, frames_per_clip=5, transform_train=batch_transform_train, transform_val=batch_transform_val, excluded_folders=[], balanced_dataset=True) ds.target_transform = get_target_transform(ds.class_to_idx, *positive_labels) print('\n\nN samples in each class:', ds.class_n_samples) print('Positive weight:', ds.pos_weight) indices = list(range(len(ds))) np.random.shuffle(indices) split_train = int(np.floor(0.7 * len(indices))) train_indices, val_indices = indices[:split_train], indices[split_train:] print('Train indices:', train_indices) print('Val indices:', val_indices) ds.val_indices = val_indices ds.set_target_to_indices_dict(ds.target_transform, ds.samples) print(ds.target_to_indices_train, ds.target_to_indices_val) train_sampler = SubsetRandomSampler(train_indices) val_sampler = SubsetRandomSampler(val_indices) print(ds.classes, ds.class_to_idx) for sample in ds.samples: print(sample) print('Creating the dataloader now') loader = DataLoader(ds, batch_size=5, sampler=train_sampler) for (x,y) in loader: print(x.size(), y) print('===') loader = DataLoader(ds, batch_size=5, sampler=val_sampler) for (x,y) in loader: print(x.size(), y)