Spaces:
Runtime error
Runtime error
| 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 | |
| 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) | |