PascalLiu/FNeVR_demo
1
1"""
2Code from https://github.com/hassony2/torch_videovision
3"""
4
5import numbers
6
7import random
8import numpy as np
9import PIL
10
11from skimage.transform import resize, rotate
12from skimage.util import pad
13import torchvision
14
15import warnings
16
17from skimage import img_as_ubyte, img_as_float
18
19
20def crop_clip(clip, min_h, min_w, h, w):
21 if isinstance(clip[0], np.ndarray):
22 cropped = [img[min_h:min_h + h, min_w:min_w + w, :] for img in clip]
23
24 elif isinstance(clip[0], PIL.Image.Image):
25 cropped = [
26 img.crop((min_w, min_h, min_w + w, min_h + h)) for img in clip
27 ]
28 else:
29 raise TypeError('Expected numpy.ndarray or PIL.Image' +
30 'but got list of {0}'.format(type(clip[0])))
31 return cropped
32
33
34def pad_clip(clip, h, w):
35 im_h, im_w = clip[0].shape[:2]
36 pad_h = (0, 0) if h < im_h else ((h - im_h) // 2, (h - im_h + 1) // 2)
37 pad_w = (0, 0) if w < im_w else ((w - im_w) // 2, (w - im_w + 1) // 2)
38
39 return pad(clip, ((0, 0), pad_h, pad_w, (0, 0)), mode='edge')
40
41
42def resize_clip(clip, size, interpolation='bilinear'):
43 if isinstance(clip[0], np.ndarray):
44 if isinstance(size, numbers.Number):
45 im_h, im_w, im_c = clip[0].shape
46 # Min spatial dim already matches minimal size
47 if (im_w <= im_h and im_w == size) or (im_h <= im_w
48 and im_h == size):
49 return clip
50 new_h, new_w = get_resize_sizes(im_h, im_w, size)
51 size = (new_w, new_h)
52 else:
53 size = size[1], size[0]
54
55 scaled = [
56 resize(img, size, order=1 if interpolation == 'bilinear' else 0, preserve_range=True,
57 mode='constant', anti_aliasing=True) for img in clip
58 ]
59 elif isinstance(clip[0], PIL.Image.Image):
60 if isinstance(size, numbers.Number):
61 im_w, im_h = clip[0].size
62 # Min spatial dim already matches minimal size
63 if (im_w <= im_h and im_w == size) or (im_h <= im_w
64 and im_h == size):
65 return clip
66 new_h, new_w = get_resize_sizes(im_h, im_w, size)
67 size = (new_w, new_h)
68 else:
69 size = size[1], size[0]
70 if interpolation == 'bilinear':
71 pil_inter = PIL.Image.NEAREST
72 else:
73 pil_inter = PIL.Image.BILINEAR
74 scaled = [img.resize(size, pil_inter) for img in clip]
75 else:
76 raise TypeError('Expected numpy.ndarray or PIL.Image' +
77 'but got list of {0}'.format(type(clip[0])))
78 return scaled
79
80
81def get_resize_sizes(im_h, im_w, size):
82 if im_w < im_h:
83 ow = size
84 oh = int(size * im_h / im_w)
85 else:
86 oh = size
87 ow = int(size * im_w / im_h)
88 return oh, ow
89
90
91class RandomFlip(object):
92 def __init__(self, time_flip=False, horizontal_flip=False):
93 self.time_flip = time_flip
94 self.horizontal_flip = horizontal_flip
95
96 def __call__(self, clip):
97 if random.random() < 0.5 and self.time_flip:
98 return clip[::-1]
99 if random.random() < 0.5 and self.horizontal_flip:
100 return [np.fliplr(img) for img in clip]
101
102 return clip
103
104
105class RandomResize(object):
106 """Resizes a list of (H x W x C) numpy.ndarray to the final size
107 The larger the original image is, the more times it takes to
108 interpolate
109 Args:
110 interpolation (str): Can be one of 'nearest', 'bilinear'
111 defaults to nearest
112 size (tuple): (widht, height)
113 """
114
115 def __init__(self, ratio=(3. / 4., 4. / 3.), interpolation='nearest'):
116 self.ratio = ratio
117 self.interpolation = interpolation
118
119 def __call__(self, clip):
120 scaling_factor = random.uniform(self.ratio[0], self.ratio[1])
121
122 if isinstance(clip[0], np.ndarray):
123 im_h, im_w, im_c = clip[0].shape
124 elif isinstance(clip[0], PIL.Image.Image):
125 im_w, im_h = clip[0].size
126
127 new_w = int(im_w * scaling_factor)
128 new_h = int(im_h * scaling_factor)
129 new_size = (new_w, new_h)
130 resized = resize_clip(
131 clip, new_size, interpolation=self.interpolation)
132
133 return resized
134
135
136class RandomCrop(object):
137 """Extract random crop at the same location for a list of videos
138 Args:
139 size (sequence or int): Desired output size for the
140 crop in format (h, w)
141 """
142
143 def __init__(self, size):
144 if isinstance(size, numbers.Number):
145 size = (size, size)
146
147 self.size = size
148
149 def __call__(self, clip):
150 """
151 Args:
152 img (PIL.Image or numpy.ndarray): List of videos to be cropped
153 in format (h, w, c) in numpy.ndarray
154 Returns:
155 PIL.Image or numpy.ndarray: Cropped list of videos
156 """
157 h, w = self.size
158 if isinstance(clip[0], np.ndarray):
159 im_h, im_w, im_c = clip[0].shape
160 elif isinstance(clip[0], PIL.Image.Image):
161 im_w, im_h = clip[0].size
162 else:
163 raise TypeError('Expected numpy.ndarray or PIL.Image' +
164 'but got list of {0}'.format(type(clip[0])))
165
166 clip = pad_clip(clip, h, w)
167 im_h, im_w = clip.shape[1:3]
168 x1 = 0 if h == im_h else random.randint(0, im_w - w)
169 y1 = 0 if w == im_w else random.randint(0, im_h - h)
170 cropped = crop_clip(clip, y1, x1, h, w)
171
172 return cropped
173
174
175class RandomRotation(object):
176 """Rotate entire clip randomly by a random angle within
177 given bounds
178 Args:
179 degrees (sequence or int): Range of degrees to select from
180 If degrees is a number instead of sequence like (min, max),
181 the range of degrees, will be (-degrees, +degrees).
182 """
183
184 def __init__(self, degrees):
185 if isinstance(degrees, numbers.Number):
186 if degrees < 0:
187 raise ValueError('If degrees is a single number,'
188 'must be positive')
189 degrees = (-degrees, degrees)
190 else:
191 if len(degrees) != 2:
192 raise ValueError('If degrees is a sequence,'
193 'it must be of len 2.')
194
195 self.degrees = degrees
196
197 def __call__(self, clip):
198 """
199 Args:
200 img (PIL.Image or numpy.ndarray): List of videos to be cropped
201 in format (h, w, c) in numpy.ndarray
202 Returns:
203 PIL.Image or numpy.ndarray: Cropped list of videos
204 """
205 angle = random.uniform(self.degrees[0], self.degrees[1])
206 if isinstance(clip[0], np.ndarray):
207 rotated = [rotate(image=img, angle=angle, preserve_range=True) for img in clip]
208 elif isinstance(clip[0], PIL.Image.Image):
209 rotated = [img.rotate(angle) for img in clip]
210 else:
211 raise TypeError('Expected numpy.ndarray or PIL.Image' +
212 'but got list of {0}'.format(type(clip[0])))
213
214 return rotated
215
216
217class ColorJitter(object):
218 """Randomly change the brightness, contrast and saturation and hue of the clip
219 Args:
220 brightness (float): How much to jitter brightness. brightness_factor
221 is chosen uniformly from [max(0, 1 - brightness), 1 + brightness].
222 contrast (float): How much to jitter contrast. contrast_factor
223 is chosen uniformly from [max(0, 1 - contrast), 1 + contrast].
224 saturation (float): How much to jitter saturation. saturation_factor
225 is chosen uniformly from [max(0, 1 - saturation), 1 + saturation].
226 hue(float): How much to jitter hue. hue_factor is chosen uniformly from
227 [-hue, hue]. Should be >=0 and <= 0.5.
228 """
229
230 def __init__(self, brightness=0, contrast=0, saturation=0, hue=0):
231 self.brightness = brightness
232 self.contrast = contrast
233 self.saturation = saturation
234 self.hue = hue
235
236 def get_params(self, brightness, contrast, saturation, hue):
237 if brightness > 0:
238 brightness_factor = random.uniform(
239 max(0, 1 - brightness), 1 + brightness)
240 else:
241 brightness_factor = None
242
243 if contrast > 0:
244 contrast_factor = random.uniform(
245 max(0, 1 - contrast), 1 + contrast)
246 else:
247 contrast_factor = None
248
249 if saturation > 0:
250 saturation_factor = random.uniform(
251 max(0, 1 - saturation), 1 + saturation)
252 else:
253 saturation_factor = None
254
255 if hue > 0:
256 hue_factor = random.uniform(-hue, hue)
257 else:
258 hue_factor = None
259 return brightness_factor, contrast_factor, saturation_factor, hue_factor
260
261 def __call__(self, clip):
262 """
263 Args:
264 clip (list): list of PIL.Image
265 Returns:
266 list PIL.Image : list of transformed PIL.Image
267 """
268 if isinstance(clip[0], np.ndarray):
269 brightness, contrast, saturation, hue = self.get_params(
270 self.brightness, self.contrast, self.saturation, self.hue)
271
272 # Create img transform function sequence
273 img_transforms = []
274 if brightness is not None:
275 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_brightness(img, brightness))
276 if saturation is not None:
277 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_saturation(img, saturation))
278 if hue is not None:
279 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_hue(img, hue))
280 if contrast is not None:
281 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_contrast(img, contrast))
282 random.shuffle(img_transforms)
283 img_transforms = [img_as_ubyte, torchvision.transforms.ToPILImage()] + img_transforms + [np.array,
284 img_as_float]
285
286 with warnings.catch_warnings():
287 warnings.simplefilter("ignore")
288 jittered_clip = []
289 for img in clip:
290 jittered_img = img
291 for func in img_transforms:
292 jittered_img = func(jittered_img)
293 jittered_clip.append(jittered_img.astype('float32'))
294 elif isinstance(clip[0], PIL.Image.Image):
295 brightness, contrast, saturation, hue = self.get_params(
296 self.brightness, self.contrast, self.saturation, self.hue)
297
298 # Create img transform function sequence
299 img_transforms = []
300 if brightness is not None:
301 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_brightness(img, brightness))
302 if saturation is not None:
303 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_saturation(img, saturation))
304 if hue is not None:
305 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_hue(img, hue))
306 if contrast is not None:
307 img_transforms.append(lambda img: torchvision.transforms.functional.adjust_contrast(img, contrast))
308 random.shuffle(img_transforms)
309
310 # Apply to all videos
311 jittered_clip = []
312 for img in clip:
313 for func in img_transforms:
314 jittered_img = func(img)
315 jittered_clip.append(jittered_img)
316
317 else:
318 raise TypeError('Expected numpy.ndarray or PIL.Image' +
319 'but got list of {0}'.format(type(clip[0])))
320 return jittered_clip
321
322
323class AllAugmentationTransform:
324 def __init__(self, resize_param=None, rotation_param=None, flip_param=None, crop_param=None, jitter_param=None):
325 self.transforms = []
326
327 if flip_param is not None:
328 self.transforms.append(RandomFlip(**flip_param))
329
330 if rotation_param is not None:
331 self.transforms.append(RandomRotation(**rotation_param))
332
333 if resize_param is not None:
334 self.transforms.append(RandomResize(**resize_param))
335
336 if crop_param is not None:
337 self.transforms.append(RandomCrop(**crop_param))
338
339 if jitter_param is not None:
340 self.transforms.append(ColorJitter(**jitter_param))
341
342 def __call__(self, clip):
343 for t in self.transforms:
344 clip = t(clip)
345 return clip
346 