Team Ai
Apppublic

PascalLiu/FNeVR_demo

sourceHugging Faceapache-2.0updated 4y agoView on Hugging Face
1likes
augmentation.py346 linesDownload Raw Back to root
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