AnsenH's picture
feat: add our model
24615d9
Raw
History Blame Contribute Delete
8.55 kB
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)