new project from v0.0.1
This commit is contained in:
@@ -0,0 +1,26 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.transform.augmention import ColorJitterGeneral
|
||||
from scepter.modules.transform.compose import Compose
|
||||
from scepter.modules.transform.identity import Identity
|
||||
from scepter.modules.transform.image import (CenterCrop, FlexibleCenterCrop,
|
||||
FlexibleResize, ImageToTensor,
|
||||
ImageTransform, Normalize,
|
||||
RandomHorizontalFlip,
|
||||
RandomResizedCrop, Resize)
|
||||
from scepter.modules.transform.io import (LoadCvImageFromFile,
|
||||
LoadImageFromFile,
|
||||
LoadImageFromFileList,
|
||||
LoadPILImageFromFile)
|
||||
from scepter.modules.transform.io_video import (DecodeVideoToTensor,
|
||||
LoadVideoFromFile)
|
||||
from scepter.modules.transform.registry import TRANSFORMS, build_pipeline
|
||||
from scepter.modules.transform.tensor import Rename, Select, ToTensor
|
||||
from scepter.modules.transform.transform_xl import FlexibleCropXL
|
||||
from scepter.modules.transform.video import (AutoResizedCropVideo,
|
||||
CenterCropVideo, NormalizeVideo,
|
||||
RandomHorizontalFlipVideo,
|
||||
RandomResizedCropVideo,
|
||||
ResizeVideo, VideoToTensor,
|
||||
VideoTransform)
|
||||
@@ -0,0 +1,506 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import numbers
|
||||
import random
|
||||
|
||||
# https://github.com/TengdaHan/DPC/blob/master/utils/augmentation.py
|
||||
import torch
|
||||
from torchvision.transforms import Compose, Lambda
|
||||
|
||||
from scepter.modules.transform.image import ImageTransform
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.transform.utils import (BACKEND_CV2, BACKEND_PILLOW,
|
||||
BACKEND_TORCHVISION,
|
||||
TORCHVISION_CAPABILITY)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
if TORCHVISION_CAPABILITY:
|
||||
BACKENDS = (BACKEND_PILLOW, BACKEND_CV2, BACKEND_TORCHVISION)
|
||||
else:
|
||||
BACKENDS = (BACKEND_PILLOW, BACKEND_CV2)
|
||||
|
||||
|
||||
def _is_tensor_a_torch_image(input):
|
||||
return input.ndim >= 2
|
||||
|
||||
|
||||
def _blend(img1, img2, ratio):
|
||||
# type: (torch.Tensor, torch.Tensor, float) -> torch.Tensor
|
||||
bound = 1 if img1.dtype in [torch.half, torch.float32, torch.float64
|
||||
] else 255
|
||||
return (ratio * img1 + (1 - ratio) * img2).clamp(0, bound).to(img1.dtype)
|
||||
|
||||
|
||||
def rgb_to_grayscale(img, split=False):
|
||||
# type: (torch.Tensor) -> torch.Tensor
|
||||
"""Convert the given RGB Image Tensor to Grayscale.
|
||||
For RGB to Grayscale conversion, ITU-R 601-2 luma transform is performed which
|
||||
is L = R * 0.2989 + G * 0.5870 + B * 0.1140
|
||||
Args:
|
||||
img (Tensor): Image to be converted to Grayscale in the form [C, H, W].
|
||||
Returns:
|
||||
Tensor: Grayscale image.
|
||||
Args:
|
||||
clip (torch.tensor): Size is (T, H, W, C)
|
||||
Return:
|
||||
clip (torch.tensor): Size is (T, H, W, C)
|
||||
"""
|
||||
orig_dtype = img.dtype
|
||||
rgb_convert = torch.tensor([0.299, 0.587, 0.114])
|
||||
if split:
|
||||
rgb_convert *= 0
|
||||
channel = random.randint(0, 2)
|
||||
rgb_convert[channel] = 1
|
||||
|
||||
assert img.shape[0] == 3, 'First dimension need to be 3 Channels'
|
||||
if img.is_cuda:
|
||||
rgb_convert = rgb_convert.to(img.device)
|
||||
|
||||
img = img.float().permute(1, 2, 3, 0).matmul(rgb_convert).to(orig_dtype)
|
||||
return torch.stack([img, img, img], 0)
|
||||
|
||||
|
||||
def _rgb2hsv(img):
|
||||
r, g, b = img.unbind(0)
|
||||
|
||||
maxc, _ = torch.max(img, dim=0)
|
||||
minc, _ = torch.min(img, dim=0)
|
||||
|
||||
eqc = maxc == minc
|
||||
cr = maxc - minc
|
||||
s = cr / torch.where(eqc, maxc.new_ones(()), maxc)
|
||||
cr_divisor = torch.where(eqc, maxc.new_ones(()), cr)
|
||||
rc = (maxc - r) / cr_divisor
|
||||
gc = (maxc - g) / cr_divisor
|
||||
bc = (maxc - b) / cr_divisor
|
||||
|
||||
hr = (maxc == r) * (bc - gc)
|
||||
hg = ((maxc == g) & (maxc != r)) * (2.0 + rc - bc)
|
||||
hb = ((maxc != g) & (maxc != r)) * (4.0 + gc - rc)
|
||||
h = (hr + hg + hb)
|
||||
h = torch.fmod((h / 6.0 + 1.0), 1.0)
|
||||
return torch.stack((h, s, maxc))
|
||||
|
||||
|
||||
def _hsv2rgb(img):
|
||||
l = len(img.shape) # noqa
|
||||
h, s, v = img.unbind(0)
|
||||
i = torch.floor(h * 6.0)
|
||||
f = (h * 6.0) - i
|
||||
i = i.to(dtype=torch.int32)
|
||||
|
||||
p = torch.clamp((v * (1.0 - s)), 0.0, 1.0)
|
||||
q = torch.clamp((v * (1.0 - s * f)), 0.0, 1.0)
|
||||
t = torch.clamp((v * (1.0 - s * (1.0 - f))), 0.0, 1.0)
|
||||
i = i % 6
|
||||
|
||||
if l == 3: # noqa
|
||||
tmp = torch.arange(6)[:, None, None]
|
||||
elif l == 4: # noqa
|
||||
tmp = torch.arange(6)[:, None, None, None]
|
||||
|
||||
if img.is_cuda:
|
||||
tmp = tmp.to(img.device)
|
||||
|
||||
mask = i == tmp # (H, W) == (6, H, W)
|
||||
|
||||
a1 = torch.stack((v, q, p, p, t, v))
|
||||
a2 = torch.stack((t, v, v, q, p, p))
|
||||
a3 = torch.stack((p, p, t, v, v, q))
|
||||
a4 = torch.stack((a1, a2, a3)) # (3, 6, H, W)
|
||||
|
||||
if l == 3: # noqa
|
||||
return torch.einsum('ijk, xijk -> xjk', mask.to(dtype=img.dtype),
|
||||
a4) # (C, H, W)
|
||||
elif l == 4: # noqa
|
||||
return torch.einsum('itjk, xitjk -> xtjk', mask.to(dtype=img.dtype),
|
||||
a4) # (C, T, H, W)
|
||||
|
||||
|
||||
def adjust_brightness(img, brightness_factor):
|
||||
# type: (torch.Tensor, float) -> torch.Tensor
|
||||
if not _is_tensor_a_torch_image(img):
|
||||
raise TypeError('tensor is not a torch image.')
|
||||
|
||||
return _blend(img, torch.zeros_like(img), brightness_factor)
|
||||
|
||||
|
||||
def adjust_contrast(img, contrast_factor):
|
||||
# type: (torch.Tensor, float) -> torch.Tensor
|
||||
if not _is_tensor_a_torch_image(img):
|
||||
raise TypeError('tensor is not a torch image.')
|
||||
|
||||
mean = torch.mean(rgb_to_grayscale(img).to(torch.float),
|
||||
dim=(-4, -2, -1),
|
||||
keepdim=True)
|
||||
|
||||
return _blend(img, mean, contrast_factor)
|
||||
|
||||
|
||||
def adjust_saturation(img, saturation_factor):
|
||||
# type: (torch.Tensor, float) -> torch.Tensor
|
||||
if not _is_tensor_a_torch_image(img):
|
||||
raise TypeError('tensor is not a torch image.')
|
||||
|
||||
return _blend(img, rgb_to_grayscale(img), saturation_factor)
|
||||
|
||||
|
||||
def adjust_hue(img, hue_factor):
|
||||
"""Adjust hue of an image.
|
||||
The image hue is adjusted by converting the image to HSV and
|
||||
cyclically shifting the intensities in the hue channel (H).
|
||||
The image is then converted back to original image mode.
|
||||
`hue_factor` is the amount of shift in H channel and must be in the
|
||||
interval `[-0.5, 0.5]`.
|
||||
See `Hue`_ for more details.
|
||||
.. _Hue: https://en.wikipedia.org/wiki/Hue
|
||||
Args:
|
||||
img (Tensor): Image to be adjusted. Image type is either uint8 or float.
|
||||
hue_factor (float): How much to shift the hue channel. Should be in
|
||||
[-0.5, 0.5]. 0.5 and -0.5 give complete reversal of hue channel in
|
||||
HSV space in positive and negative direction respectively.
|
||||
0 means no shift. Therefore, both -0.5 and 0.5 will give an image
|
||||
with complementary colors while 0 gives the original image.
|
||||
Returns:
|
||||
Tensor: Hue adjusted image.
|
||||
"""
|
||||
if isinstance(hue_factor, float) and not (-0.5 <= hue_factor <= 0.5):
|
||||
raise ValueError(
|
||||
'hue_factor ({}) is not in [-0.5, 0.5].'.format(hue_factor))
|
||||
elif (isinstance(hue_factor, torch.Tensor)
|
||||
and not ((-0.5 <= hue_factor).sum() == hue_factor.shape[0] and
|
||||
(hue_factor <= 0.5).sum() == hue_factor.shape[0])):
|
||||
raise ValueError(
|
||||
'hue_factor ({}) is not in [-0.5, 0.5].'.format(hue_factor))
|
||||
|
||||
if not _is_tensor_a_torch_image(img):
|
||||
raise TypeError('tensor is not a torch image.')
|
||||
|
||||
orig_dtype = img.dtype
|
||||
if img.dtype == torch.uint8:
|
||||
img = img.to(dtype=torch.float32) / 255.0
|
||||
|
||||
img = _rgb2hsv(img)
|
||||
h, s, v = img.unbind(0)
|
||||
h += hue_factor
|
||||
h = h % 1.0
|
||||
img = torch.stack((h, s, v))
|
||||
img_hue_adj = _hsv2rgb(img)
|
||||
|
||||
if orig_dtype == torch.uint8:
|
||||
img_hue_adj = (img_hue_adj * 255.0).to(dtype=orig_dtype)
|
||||
return img_hue_adj
|
||||
|
||||
|
||||
# https://github.com/TengdaHan/DPC/blob/master/utils/augmentation.py
|
||||
class ColorJitter(object):
|
||||
"""Randomly change the brightness, contrast and saturation of an image.
|
||||
Args:
|
||||
brightness (float or tuple of float (min, max)): How much to jitter brightness.
|
||||
brightness_factor is chosen uniformly from [max(0, 1 - brightness), 1 + brightness]
|
||||
or the given [min, max]. Should be non negative numbers.
|
||||
contrast (float or tuple of float (min, max)): How much to jitter contrast.
|
||||
contrast_factor is chosen uniformly from [max(0, 1 - contrast), 1 + contrast]
|
||||
or the given [min, max]. Should be non negative numbers.
|
||||
saturation (float or tuple of float (min, max)): How much to jitter saturation.
|
||||
saturation_factor is chosen uniformly from [max(0, 1 - saturation), 1 + saturation]
|
||||
or the given [min, max]. Should be non negative numbers.
|
||||
hue (float or tuple of float (min, max)): How much to jitter hue.
|
||||
hue_factor is chosen uniformly from [-hue, hue] or the given [min, max].
|
||||
Should have 0<= hue <= 0.5 or -0.5 <= min <= max <= 0.5.
|
||||
grayscale (probablitities for rgb-to-gray 0~1)
|
||||
consistent (for video input whether the augment scale is consistent or not!)
|
||||
shuffle (shuffle the transform's order when there are multiple transform)
|
||||
gray_first (whether use grayscale or not)
|
||||
is_split (whether randomly chance the channel as the gray results.)
|
||||
"""
|
||||
def __init__(self,
|
||||
brightness=0,
|
||||
contrast=0,
|
||||
saturation=0,
|
||||
hue=0,
|
||||
grayscale=0,
|
||||
consistent=False,
|
||||
shuffle=True,
|
||||
gray_first=True,
|
||||
is_split=False):
|
||||
self.brightness = self._check_input(brightness, 'brightness')
|
||||
self.contrast = self._check_input(contrast, 'contrast')
|
||||
self.saturation = self._check_input(saturation, 'saturation')
|
||||
self.hue = self._check_input(hue,
|
||||
'hue',
|
||||
center=0,
|
||||
bound=(-0.5, 0.5),
|
||||
clip_first_on_zero=False)
|
||||
self.grayscale = grayscale
|
||||
self.consistent = consistent
|
||||
self.shuffle = shuffle
|
||||
self.gray_first = gray_first
|
||||
self.is_split = is_split
|
||||
|
||||
def _check_input(self,
|
||||
value,
|
||||
name,
|
||||
center=1,
|
||||
bound=(0, float('inf')),
|
||||
clip_first_on_zero=True):
|
||||
if isinstance(value, numbers.Number):
|
||||
if value < 0:
|
||||
raise ValueError(
|
||||
'If {} is a single number, it must be non negative.'.
|
||||
format(name))
|
||||
value = [center - float(value), center + float(value)]
|
||||
if clip_first_on_zero:
|
||||
value[0] = max(value[0], 0.0)
|
||||
elif isinstance(value, (tuple, list)) and len(value) == 2:
|
||||
if not bound[0] <= value[0] <= value[1] <= bound[1]:
|
||||
raise ValueError('{} values should be between {}'.format(
|
||||
name, bound))
|
||||
else:
|
||||
raise TypeError(
|
||||
'{} should be a single number or a list/tuple with lenght 2.'.
|
||||
format(name))
|
||||
|
||||
# if value is 0 or (1., 1.) for brightness/contrast/saturation
|
||||
# or (0., 0.) for hue, do nothing
|
||||
if value[0] == value[1] == center:
|
||||
value = None
|
||||
return value
|
||||
|
||||
def _get_transform(self, T, device):
|
||||
"""Get a randomized transform to be applied on image.
|
||||
Arguments are same as that of __init__.
|
||||
Arg:
|
||||
T (int): number of frames. Used when consistent = False.
|
||||
Returns:
|
||||
Transform which randomly adjusts brightness, contrast and
|
||||
saturation in a random order.
|
||||
"""
|
||||
transforms = []
|
||||
if self.brightness is not None:
|
||||
if self.consistent:
|
||||
brightness_factor = random.uniform(self.brightness[0],
|
||||
self.brightness[1])
|
||||
else:
|
||||
brightness_factor = torch.empty([1, T, 1, 1],
|
||||
device=device).uniform_(
|
||||
self.brightness[0],
|
||||
self.brightness[1])
|
||||
transforms.append(
|
||||
Lambda(
|
||||
lambda frame: adjust_brightness(frame, brightness_factor)))
|
||||
|
||||
if self.contrast is not None:
|
||||
if self.consistent:
|
||||
contrast_factor = random.uniform(self.contrast[0],
|
||||
self.contrast[1])
|
||||
else:
|
||||
contrast_factor = torch.empty([1, T, 1, 1],
|
||||
device=device).uniform_(
|
||||
self.contrast[0],
|
||||
self.contrast[1])
|
||||
transforms.append(
|
||||
Lambda(lambda frame: adjust_contrast(frame, contrast_factor)))
|
||||
|
||||
if self.saturation is not None:
|
||||
if self.consistent:
|
||||
saturation_factor = random.uniform(self.saturation[0],
|
||||
self.saturation[1])
|
||||
else:
|
||||
saturation_factor = torch.empty([1, T, 1, 1],
|
||||
device=device).uniform_(
|
||||
self.saturation[0],
|
||||
self.saturation[1])
|
||||
transforms.append(
|
||||
Lambda(
|
||||
lambda frame: adjust_saturation(frame, saturation_factor)))
|
||||
|
||||
if self.hue is not None:
|
||||
if self.consistent:
|
||||
hue_factor = random.uniform(self.hue[0], self.hue[1])
|
||||
else:
|
||||
hue_factor = torch.empty([T, 1, 1], device=device).uniform_(
|
||||
self.hue[0], self.hue[1])
|
||||
transforms.append(
|
||||
Lambda(lambda frame: adjust_hue(frame, hue_factor)))
|
||||
|
||||
if self.shuffle:
|
||||
random.shuffle(transforms)
|
||||
|
||||
if random.uniform(0, 1) < self.grayscale:
|
||||
gray_transform = Lambda(
|
||||
lambda frame: rgb_to_grayscale(frame, split=self.is_split))
|
||||
if self.gray_first:
|
||||
transforms.insert(0, gray_transform)
|
||||
else:
|
||||
transforms.append(gray_transform)
|
||||
|
||||
transform = Compose(transforms)
|
||||
return transform
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Size is (C, T, H, W)
|
||||
Return:
|
||||
clip (torch.tensor): Size is (C, T, H, W)
|
||||
"""
|
||||
|
||||
is_frame = False
|
||||
|
||||
if len(clip.shape) == 3:
|
||||
# frame data, transfer to 4dim
|
||||
clip = torch.unsqueeze(clip, dim=1)
|
||||
is_frame = True
|
||||
# (C, T, H, W)
|
||||
raw_shape = clip.shape
|
||||
device = clip.device
|
||||
T = raw_shape[1]
|
||||
transform = self._get_transform(T, device)
|
||||
clip = transform(clip)
|
||||
assert clip.shape == raw_shape
|
||||
if is_frame:
|
||||
clip = torch.squeeze(clip)
|
||||
return clip # (C, T, H, W)
|
||||
|
||||
def __repr__(self):
|
||||
format_string = self.__class__.__name__ + '('
|
||||
format_string += 'brightness={0}'.format(self.brightness)
|
||||
format_string += ', contrast={0}'.format(self.contrast)
|
||||
format_string += ', saturation={0}'.format(self.saturation)
|
||||
format_string += ', hue={0})'.format(self.hue)
|
||||
format_string += ', grayscale={0})'.format(self.grayscale)
|
||||
return format_string
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class ColorJitterGeneral(ImageTransform):
|
||||
'''
|
||||
brightness (float or tuple of float (min, max)): How much to jitter brightness.
|
||||
brightness_factor is chosen uniformly from [max(0, 1 - brightness), 1 + brightness]
|
||||
or the given [min, max]. Should be non negative numbers.
|
||||
contrast (float or tuple of float (min, max)): How much to jitter contrast.
|
||||
contrast_factor is chosen uniformly from [max(0, 1 - contrast), 1 + contrast]
|
||||
or the given [min, max]. Should be non negative numbers.
|
||||
saturation (float or tuple of float (min, max)): How much to jitter saturation.
|
||||
saturation_factor is chosen uniformly from [max(0, 1 - saturation), 1 + saturation]
|
||||
or the given [min, max]. Should be non negative numbers.
|
||||
hue (float or tuple of float (min, max)): How much to jitter hue.
|
||||
hue_factor is chosen uniformly from [-hue, hue] or the given [min, max].
|
||||
Should have 0<= hue <= 0.5 or -0.5 <= min <= max <= 0.5.
|
||||
grayscale (probablitities for rgb-to-gray 0~1)
|
||||
consistent (for video input whether the augment scale is consistent or not!)
|
||||
shuffle (shuffle the transform's order when there are multiple transform)
|
||||
gray_first (whether use grayscale or not)
|
||||
is_split (whether randomly chance the channel as the gray results.)
|
||||
'''
|
||||
para_dict = [{
|
||||
'BRIGHTNESS': {
|
||||
'value':
|
||||
0,
|
||||
'description':
|
||||
'(float or tuple of float (min, max)): How much to jitter brightness.'
|
||||
'brightness_factor is chosen uniformly from [max(0, 1 - brightness), 1 + brightness]'
|
||||
'or the given [min, max]. Should be non negative numbers.'
|
||||
},
|
||||
'CONTRAST': {
|
||||
'value':
|
||||
0,
|
||||
'description':
|
||||
'(float or tuple of float (min, max)): How much to jitter contrast.'
|
||||
'contrast_factor is chosen uniformly from [max(0, 1 - contrast), 1 + contrast]'
|
||||
'or the given [min, max]. Should be non negative numbers.'
|
||||
},
|
||||
'SATURATION': {
|
||||
'value':
|
||||
0,
|
||||
'description':
|
||||
'(float or tuple of float (min, max)): How much to jitter saturation.'
|
||||
'saturation_factor is chosen uniformly from [max(0, 1 - saturation), 1 + saturation]'
|
||||
'or the given [min, max]. Should be non negative numbers.'
|
||||
},
|
||||
'HUE': {
|
||||
'value':
|
||||
0,
|
||||
'description':
|
||||
'(float or tuple of float (min, max)): How much to jitter hue.'
|
||||
'hue_factor is chosen uniformly from [-hue, hue] or the given [min, max].'
|
||||
'Should have 0<= hue <= 0.5 or -0.5 <= min <= max <= 0.5.'
|
||||
},
|
||||
'GRAYSCALE': {
|
||||
'value': 0,
|
||||
'description': '(probablitities for rgb-to-gray 0~1)'
|
||||
},
|
||||
'CONSISTENT': {
|
||||
'value':
|
||||
False,
|
||||
'description':
|
||||
'(for video input whether the augment scale is consistent or not!)'
|
||||
},
|
||||
'SHUFFLE': {
|
||||
'value':
|
||||
False,
|
||||
'description':
|
||||
"(shuffle the transform's order when there are multiple transform)"
|
||||
},
|
||||
'GRAY_FIRST': {
|
||||
'value': False,
|
||||
'description': '(whether use grayscale or not)'
|
||||
},
|
||||
'IS_SPLIT': {
|
||||
'value':
|
||||
False,
|
||||
'description':
|
||||
'(whether randomly chance the channel as the gray results.)'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
para_dict[0].pop('BACKEND')
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(ColorJitterGeneral, self).__init__(cfg, logger=None)
|
||||
# brightness=0, contrast=0, saturation=0, hue=0, grayscale=0,
|
||||
# consistent=False, shuffle=True, gray_first=True, is_split=False
|
||||
brightness = cfg.get('BRIGHTNESS', 0)
|
||||
contrast = cfg.get('CONTRAST', 0)
|
||||
saturation = cfg.get('SATURATION', 0)
|
||||
hue = cfg.get('HUE', 0)
|
||||
grayscale = cfg.get('GRAYSCALE', 0)
|
||||
consistent = cfg.get('CONSISTENT', False)
|
||||
shuffle = cfg.get('SHUFFLE', False)
|
||||
gray_first = cfg.get('GRAY_FIRST', False)
|
||||
is_split = cfg.get('IS_SPLIT', False)
|
||||
|
||||
cj_ins = ColorJitter(brightness=brightness,
|
||||
contrast=contrast,
|
||||
saturation=saturation,
|
||||
hue=hue,
|
||||
grayscale=grayscale,
|
||||
consistent=consistent,
|
||||
shuffle=shuffle,
|
||||
gray_first=gray_first,
|
||||
is_split=is_split)
|
||||
self.callable = cj_ins
|
||||
|
||||
def __call__(self, item):
|
||||
item[self.output_key] = self.callable(item[self.input_key])
|
||||
# item["meta"]["normalize_params"] = dict(mean=self.mean, std=self.std)
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
ColorJitterGeneral.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,39 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class Compose(object):
|
||||
""" Compose all transform function into one.
|
||||
|
||||
Args:
|
||||
transform (List[dict]): List of transform configs.
|
||||
|
||||
"""
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.transforms = [TRANSFORMS.build(t) for t in cfg.TRANSFORMS]
|
||||
|
||||
def __call__(self, item):
|
||||
for t in self.transforms:
|
||||
item = t(item)
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__, [{}],
|
||||
set_name=True)
|
||||
@@ -0,0 +1,31 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from ..utils.config import dict_to_yaml
|
||||
from .registry import TRANSFORMS
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class Identity(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
pass
|
||||
|
||||
def __call__(self, item):
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__, [{}],
|
||||
set_name=True)
|
||||
@@ -0,0 +1,626 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import numbers
|
||||
|
||||
import numpy as np
|
||||
import opencv_transforms.functional as cv2_TF
|
||||
import opencv_transforms.transforms as cv2_transforms
|
||||
import torch
|
||||
import torchvision.transforms as transforms
|
||||
import torchvision.transforms.functional as TF
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.transform.utils import (
|
||||
BACKEND_CV2, BACKEND_PILLOW, BACKEND_TORCHVISION, INPUT_CV2_TYPE_WARNING,
|
||||
INPUT_PIL_TYPE_WARNING, INPUT_TENSOR_TYPE_WARNING, INTERPOLATION_STYLE,
|
||||
INTERPOLATION_STYLE_CV2, TORCHVISION_CAPABILITY, is_cv2_image,
|
||||
is_pil_image, is_tensor)
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
if TORCHVISION_CAPABILITY:
|
||||
BACKENDS = (BACKEND_PILLOW, BACKEND_CV2, BACKEND_TORCHVISION)
|
||||
else:
|
||||
BACKENDS = (BACKEND_PILLOW, BACKEND_CV2)
|
||||
|
||||
|
||||
class ImageTransform(object):
|
||||
para_dict = [{
|
||||
'INPUT_KEY': {
|
||||
'value': 'img',
|
||||
'description': 'input key or key list.'
|
||||
},
|
||||
'OUTPUT_KEY': {
|
||||
'value': 'img',
|
||||
'description': 'input key or key list.'
|
||||
},
|
||||
'BACKEND': {
|
||||
'value': 'pillow',
|
||||
'description': 'backend, choose from pillow, cv2, torchvision'
|
||||
}
|
||||
}]
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.input_key = cfg.get('INPUT_KEY', 'img')
|
||||
self.output_key = cfg.get('OUTPUT_KEY', 'img')
|
||||
self.backend = cfg.get('BACKEND', BACKEND_PILLOW)
|
||||
|
||||
def check_image_type(self, input_img):
|
||||
if self.backend == BACKEND_PILLOW:
|
||||
assert is_pil_image(input_img), INPUT_PIL_TYPE_WARNING
|
||||
w, h = input_img.size
|
||||
return h, w
|
||||
elif self.backend == BACKEND_CV2:
|
||||
assert is_cv2_image(input_img), INPUT_CV2_TYPE_WARNING
|
||||
h, w, c = input_img.shape
|
||||
return h, w
|
||||
elif TORCHVISION_CAPABILITY:
|
||||
if self.backend == BACKEND_TORCHVISION:
|
||||
assert is_tensor(input_img), INPUT_TENSOR_TYPE_WARNING
|
||||
c, h, w = input_img.shape
|
||||
return h, w
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
ImageTransform.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class RandomCrop(ImageTransform):
|
||||
""" Crop a random portion of image.
|
||||
If the image is torch Tensor, it is expected to have [..., H, W] shape.
|
||||
|
||||
Args:
|
||||
size (sequence or int): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the output size will be matched to (size, size).
|
||||
padding (sequence or int): Optional padding on each border of the image. Default is None.
|
||||
pad_if_needed (bool): It will pad the image if smaller than the desired size to avoid raising an exception.
|
||||
fill (number or str or tuple): Pixel fill value for constant fill. Default is 0.
|
||||
padding_mode (str): Type of padding. Should be: constant, edge, reflect or symmetric.
|
||||
Default is constant.
|
||||
"""
|
||||
para_dict = [{
|
||||
'SIZE': {
|
||||
'value': 224,
|
||||
'description': 'crop size'
|
||||
},
|
||||
'PADDING': {
|
||||
'value': None,
|
||||
'description': 'padding'
|
||||
},
|
||||
'PAD_IF_NEEDED': {
|
||||
'value': False,
|
||||
'description': 'pad if needed'
|
||||
},
|
||||
'FILL': {
|
||||
'value': 0,
|
||||
'description': 'fill'
|
||||
},
|
||||
'PADDING_MODE': {
|
||||
'value': 'constant',
|
||||
'description': 'padding mode'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
size = cfg.SIZE
|
||||
padding = cfg.get('PADDING', None)
|
||||
pad_if_needed = cfg.get('PAD_IF_NEEDED', False)
|
||||
fill = cfg.get('FILL', 0)
|
||||
padding_mode = cfg.get('PADDING_MODE', 'constant')
|
||||
super(RandomCrop, self).__init__(cfg)
|
||||
assert self.backend in BACKENDS
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION):
|
||||
self.callable = transforms.RandomCrop(size,
|
||||
padding=padding,
|
||||
pad_if_needed=pad_if_needed,
|
||||
fill=fill,
|
||||
padding_mode=padding_mode)
|
||||
else:
|
||||
self.callable = cv2_transforms.RandomCrop(
|
||||
size,
|
||||
padding=padding,
|
||||
pad_if_needed=pad_if_needed,
|
||||
fill=fill,
|
||||
padding_mode=padding_mode)
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
for idx, key in enumerate(self.input_key):
|
||||
self.check_image_type(item[key])
|
||||
item[self.output_key[idx]] = self.callable(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
RandomCrop.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class RandomResizedCrop(ImageTransform):
|
||||
"""Crop a random portion of image and resize it to a given size.
|
||||
|
||||
If the image is torch Tensor, it is expected to have [..., H, W] shape.
|
||||
|
||||
Args:
|
||||
size (int or sequence): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the output size will be matched to (size, size).
|
||||
scale (tuple of float): Specifies the lower and upper bounds for the random area of the crop,
|
||||
before resizing. The scale is defined with respect to the area of the original image.
|
||||
ratio (tuple of float): lower and upper bounds for the random aspect ratio of the crop, before
|
||||
resizing.
|
||||
interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported.
|
||||
"""
|
||||
para_dict = [{
|
||||
'SIZE': {
|
||||
'value': 224,
|
||||
'description': 'crop size'
|
||||
},
|
||||
'RATIO': {
|
||||
'value': [3. / 4., 4. / 3.],
|
||||
'description': 'ratio'
|
||||
},
|
||||
'SCALE': {
|
||||
'value': [0.08, 1.0],
|
||||
'description': 'scale'
|
||||
},
|
||||
'INTERPOLATION': {
|
||||
'value': 'bilinear',
|
||||
'description': 'interpolation'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(RandomResizedCrop, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
self.interpolation = cfg.get('INTERPOLATION', 'bilinear')
|
||||
self.size = cfg.SIZE
|
||||
self.scale = tuple(cfg.get('SCALE', [0.08, 1.0]))
|
||||
self.ratio = tuple(cfg.get('RATIO', [3. / 4., 4. / 3.]))
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION):
|
||||
assert self.interpolation in INTERPOLATION_STYLE
|
||||
else:
|
||||
assert self.interpolation in INTERPOLATION_STYLE_CV2
|
||||
self.callable = transforms.RandomResizedCrop(self.size,
|
||||
self.scale, self.ratio, INTERPOLATION_STYLE[self.interpolation]) \
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) \
|
||||
else cv2_transforms.RandomResizedCrop(self.size, self.scale, self.ratio, INTERPOLATION_STYLE_CV2[self.interpolation]) # noqa
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
for idx, key in enumerate(self.input_key):
|
||||
self.check_image_type(item[key])
|
||||
item[self.output_key[idx]] = self.callable(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
RandomResizedCrop.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class Resize(ImageTransform):
|
||||
"""Resize image to a given size.
|
||||
|
||||
If the image is torch Tensor, it is expected to have [..., H, W] shape.
|
||||
|
||||
Args:
|
||||
size (int or sequence): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the smaller edge of the image will be matched to this number
|
||||
maintaining the aspect ratio.
|
||||
interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported.
|
||||
"""
|
||||
para_dict = [{
|
||||
'INTERPOLATION': {
|
||||
'value': 'bilinear',
|
||||
'description': 'interpolation'
|
||||
},
|
||||
'SIZE': {
|
||||
'value': 224,
|
||||
'description': 'resize to size 224'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(Resize, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
self.size = cfg.SIZE
|
||||
self.interpolation = cfg.get('INTERPOLATION', 'bilinear')
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION):
|
||||
assert self.interpolation in INTERPOLATION_STYLE
|
||||
else:
|
||||
assert self.interpolation in INTERPOLATION_STYLE_CV2
|
||||
self.callable = transforms.Resize(self.size, INTERPOLATION_STYLE[self.interpolation]) \
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) \
|
||||
else cv2_transforms.Resize(self.size, INTERPOLATION_STYLE_CV2[self.interpolation])
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
for idx, key in enumerate(self.input_key):
|
||||
self.check_image_type(item[key])
|
||||
item[self.output_key[idx]] = self.callable(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
Resize.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class CenterCrop(ImageTransform):
|
||||
""" Crops the given image at the center.
|
||||
|
||||
If the image is torch Tensor, it is expected to have [..., H, W] shape.
|
||||
|
||||
Args:
|
||||
size (sequence or int): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the output size will be matched to (size, size).
|
||||
"""
|
||||
para_dict = [{'SIZE': {'value': 224, 'description': 'resize to size 224'}}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(CenterCrop, self).__init__(cfg, logger=None)
|
||||
assert self.backend in BACKENDS
|
||||
self.size = cfg.SIZE
|
||||
self.callable = transforms.CenterCrop(self.size) \
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_transforms.CenterCrop(self.size)
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
for idx, key in enumerate(self.input_key):
|
||||
self.check_image_type(item[key])
|
||||
item[self.output_key[idx]] = self.callable(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
CenterCrop.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class RandomHorizontalFlip(ImageTransform):
|
||||
""" Horizontally flip the given image randomly with a given probability.
|
||||
|
||||
If the image is torch Tensor, it is expected to have [..., H, W] shape.
|
||||
|
||||
Args:
|
||||
p (float): probability of the image being flipped. Default value is 0.5
|
||||
"""
|
||||
para_dict = [{'P': {'value': 0.5, 'description': 'P'}}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(RandomHorizontalFlip, self).__init__(cfg, logger=None)
|
||||
p = cfg.get('P', 0.5)
|
||||
assert self.backend in BACKENDS
|
||||
self.callable = transforms.RandomHorizontalFlip(p) \
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_transforms.RandomHorizontalFlip(p)
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
for idx, key in enumerate(self.input_key):
|
||||
self.check_image_type(item[key])
|
||||
item[self.output_key[idx]] = self.callable(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
RandomHorizontalFlip.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class Normalize(ImageTransform):
|
||||
""" Normalize a tensor image with mean and standard deviation.
|
||||
This transform only support tensor image.
|
||||
|
||||
Args:
|
||||
mean (sequence): Sequence of means for each channel.
|
||||
std (sequence): Sequence of standard deviations for each channel.
|
||||
"""
|
||||
para_dict = [{
|
||||
'MEAN': {
|
||||
'value': [],
|
||||
'description': 'mean'
|
||||
},
|
||||
'STD': {
|
||||
'value': [],
|
||||
'description': 'std'
|
||||
},
|
||||
}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(Normalize, self).__init__(cfg, logger=None)
|
||||
assert self.backend in BACKENDS
|
||||
mean = cfg.MEAN
|
||||
std = cfg.STD
|
||||
self.mean = np.array(mean, dtype=np.float32)
|
||||
self.std = np.array(std, dtype=np.float32)
|
||||
self.callable = transforms.Normalize(self.mean, self.std) \
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_transforms.Normalize(self.mean, self.std)
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
for idx, key in enumerate(self.input_key):
|
||||
item[self.output_key[idx]] = self.callable(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
Normalize.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class ImageToTensor(ImageTransform):
|
||||
""" Convert a ``PIL Image`` or ``numpy.ndarray`` or uint8 type tensor to a float32 tensor,
|
||||
and scale output to [0.0, 1.0].
|
||||
"""
|
||||
para_dict = [{}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(ImageToTensor, self).__init__(cfg, logger)
|
||||
assert self.backend in BACKENDS
|
||||
|
||||
if self.backend == BACKEND_PILLOW:
|
||||
self.callable = transforms.ToTensor()
|
||||
elif self.backend == BACKEND_CV2:
|
||||
self.callable = cv2_transforms.ToTensor()
|
||||
else:
|
||||
self.callable = transforms.ConvertImageDtype(torch.float)
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
for idx, key in enumerate(self.input_key):
|
||||
item[self.output_key[idx]] = self.callable(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
ImageToTensor.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class FlexibleResize(ImageTransform):
|
||||
para_dict = [{
|
||||
'INTERPOLATION': {
|
||||
'value': 'bilinear',
|
||||
'description': 'interpolation'
|
||||
},
|
||||
}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(FlexibleResize, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
self.size = cfg.get('SIZE', None)
|
||||
if self.size is not None:
|
||||
if isinstance(self.size, numbers.Number):
|
||||
self.size = [self.size, self.size]
|
||||
|
||||
interpolation = cfg.get('INTERPOLATION', 'bilinear')
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION):
|
||||
assert interpolation in INTERPOLATION_STYLE
|
||||
else:
|
||||
assert interpolation in INTERPOLATION_STYLE_CV2
|
||||
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION):
|
||||
self.callable = TF.resize
|
||||
self.interpolation = INTERPOLATION_STYLE[interpolation]
|
||||
else:
|
||||
self.callable = cv2_TF.resize
|
||||
self.interpolation = INTERPOLATION_STYLE_CV2[interpolation]
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
|
||||
for idx, key in enumerate(self.input_key):
|
||||
ih, iw = self.check_image_type(item[key])
|
||||
meta = item.get('meta', {})
|
||||
if 'image_size' in meta:
|
||||
iw, ih, ow, oh = iw, ih, meta['image_size'][1], meta[
|
||||
'image_size'][0]
|
||||
elif self.size is not None:
|
||||
iw, ih, ow, oh = iw, ih, self.size[1], self.size[0]
|
||||
meta['image_size'] = [oh, ow]
|
||||
else:
|
||||
raise KeyError(
|
||||
'The meta of input item must consists of '
|
||||
"['width', 'height', 'image_size'], and at least one key is missing."
|
||||
)
|
||||
scale = max(ow / iw, oh / ih)
|
||||
new_size = (round(scale * ih), round(scale * iw))
|
||||
item[self.output_key[idx]] = self.callable(item[key], new_size,
|
||||
self.interpolation)
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
FlexibleResize.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class FlexibleCenterCrop(ImageTransform):
|
||||
para_dict = [{}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(FlexibleCenterCrop, self).__init__(cfg, logger=None)
|
||||
assert self.backend in BACKENDS
|
||||
self.size = cfg.get('SIZE', None)
|
||||
if self.size is not None:
|
||||
if isinstance(self.size, numbers.Number):
|
||||
self.size = [self.size, self.size]
|
||||
self.callable = TF.center_crop if self.backend in (
|
||||
BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_TF.center_crop
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
|
||||
meta = item.get('meta', {})
|
||||
if 'image_size' in meta:
|
||||
oh, ow = meta['image_size']
|
||||
out_size = (oh, ow)
|
||||
else:
|
||||
out_size = self.size
|
||||
|
||||
for idx, key in enumerate(self.input_key):
|
||||
self.check_image_type(item[key])
|
||||
item[self.output_key[idx]] = self.callable(item[key], out_size)
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
FlexibleCenterCrop.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,341 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import io
|
||||
import os
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
from scepter.modules.utils.file_system import DATA_FS as FS
|
||||
|
||||
|
||||
def pillow_convert(image, rgb_order):
|
||||
if image.mode != rgb_order:
|
||||
if image.mode == 'P':
|
||||
image = image.convert(f'{rgb_order}A')
|
||||
if image.mode == f'{rgb_order}A':
|
||||
bg = Image.new(rgb_order,
|
||||
size=(image.width, image.height),
|
||||
color=(255, 255, 255))
|
||||
bg.paste(image, (0, 0), mask=image)
|
||||
image = bg
|
||||
else:
|
||||
image = image.convert('RGB')
|
||||
return image
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class LoadImageFromFile(object):
|
||||
""" Load Image from file. We have multi ways to load image. Here we compose them into one transform.
|
||||
|
||||
Args:
|
||||
rgb_order (str): 'RGB' or 'BGR'.
|
||||
backend (str): 'pillow', 'cv2' or 'torchvision'. Image should be read as uint8 dtype.
|
||||
- 'pillow': Read image file as PIL.Image object.
|
||||
- 'cv2': Read image file as numpy.ndarray object.
|
||||
- 'torchvision': Read image file as tensor object.
|
||||
"""
|
||||
def __init__(self, cfg, logger=None):
|
||||
rgb_order = cfg.get('RGB_ORDER', 'RGB')
|
||||
backend = cfg.get('BACKEND', 'pillow')
|
||||
assert rgb_order in ('RGB', 'BGR')
|
||||
assert backend in ('pillow', 'cv2', 'torchvision')
|
||||
self.rgb_order = rgb_order
|
||||
self.backend = backend
|
||||
|
||||
def read_file(self, img_path):
|
||||
if not we.data_online:
|
||||
with FS.get_from(img_path) as img_path:
|
||||
if self.backend == 'pillow':
|
||||
try:
|
||||
image = Image.open(img_path)
|
||||
image = pillow_convert(image, self.rgb_order)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
elif self.backend == 'cv2':
|
||||
image = cv2.imread(img_path, cv2.IMREAD_COLOR)
|
||||
if self.rgb_order == 'RGB':
|
||||
try:
|
||||
cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
else:
|
||||
image = Image.open(img_path).convert(self.rgb_order)
|
||||
image_np = np.asarray(image).transpose(
|
||||
(2, 0, 1)) # Tensor type needs shape to be (C, H, W)
|
||||
image = torch.from_numpy(image_np)
|
||||
return image
|
||||
else:
|
||||
with FS.get_object(img_path) as image_data:
|
||||
if self.backend == 'pillow':
|
||||
try:
|
||||
image = Image.open(io.BytesIO(image_data))
|
||||
image = pillow_convert(image, self.rgb_order)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
elif self.backend == 'cv2':
|
||||
image = cv2.imdecode(
|
||||
np.array(bytearray(image_data), dtype='uint8'),
|
||||
cv2.IMREAD_COLOR)
|
||||
if self.rgb_order == 'RGB':
|
||||
try:
|
||||
cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
else:
|
||||
image = Image.open(img_path).convert(self.rgb_order)
|
||||
image_np = np.asarray(image).transpose(
|
||||
(2, 0, 1)) # Tensor type needs shape to be (C, H, W)
|
||||
image = torch.from_numpy(image_np)
|
||||
return image
|
||||
|
||||
def __call__(self, item):
|
||||
if 'prefix' in item['meta']:
|
||||
img_path = os.path.join(item['meta']['prefix'],
|
||||
item['meta']['img_path'])
|
||||
else:
|
||||
img_path = item['meta']['img_path']
|
||||
item['img'] = self.read_file(img_path)
|
||||
item['meta']['rgb_order'] = self.rgb_order
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'RGB_ORDER': {
|
||||
'value': 'RGB',
|
||||
'description': 'rgb order'
|
||||
},
|
||||
'BACKEND': {
|
||||
'value': 'pillow',
|
||||
'description': 'input backend'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class LoadImageFromFileList(object):
|
||||
""" Load Image from file. We have multi ways to load image. Here we compose them into one transform.
|
||||
|
||||
Args:
|
||||
rgb_order (str): 'RGB' or 'BGR'.
|
||||
backend (str): 'pillow', 'cv2' or 'torchvision'. Image should be read as uint8 dtype.
|
||||
- 'pillow': Read image file as PIL.Image object.
|
||||
- 'cv2': Read image file as numpy.ndarray object.
|
||||
- 'torchvision': Read image file as tensor object.
|
||||
"""
|
||||
para_dict = [{
|
||||
'RGB_ORDER': {
|
||||
'value': 'RGB',
|
||||
'description': 'Rgb order!'
|
||||
},
|
||||
'BACKEND': {
|
||||
'value': 'pillow',
|
||||
'description': 'Input backend!'
|
||||
},
|
||||
'FILE_KEYS': {
|
||||
'value': [],
|
||||
'description':
|
||||
"The file keys for input, if key include '_path', "
|
||||
"the return results will be saved with key as key.replace('_path', '')!"
|
||||
}
|
||||
}]
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
rgb_order = cfg.get('RGB_ORDER', 'RGB')
|
||||
backend = cfg.get('BACKEND', 'pillow')
|
||||
self.file_keys = cfg.get('FILE_KEYS', ['img_path'])
|
||||
if isinstance(self.file_keys, str):
|
||||
self.file_keys = [self.file_keys]
|
||||
assert rgb_order in ('RGB', 'BGR')
|
||||
assert backend in ('pillow', 'cv2', 'torchvision')
|
||||
self.rgb_order = rgb_order
|
||||
self.backend = backend
|
||||
|
||||
def read_file(self, img_path):
|
||||
if not we.data_online:
|
||||
with FS.get_from(img_path) as img_path:
|
||||
if self.backend == 'pillow':
|
||||
try:
|
||||
image = Image.open(img_path)
|
||||
image = pillow_convert(image, self.rgb_order)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
elif self.backend == 'cv2':
|
||||
image = cv2.imread(img_path, cv2.IMREAD_COLOR)
|
||||
if self.rgb_order == 'RGB':
|
||||
try:
|
||||
cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
else:
|
||||
image = Image.open(img_path).convert(self.rgb_order)
|
||||
image_np = np.asarray(image).transpose(
|
||||
(2, 0, 1)) # Tensor type needs shape to be (C, H, W)
|
||||
image = torch.from_numpy(image_np)
|
||||
return image
|
||||
else:
|
||||
with FS.get_object(img_path) as image_data:
|
||||
if self.backend == 'pillow':
|
||||
try:
|
||||
image = Image.open(io.BytesIO(image_data))
|
||||
image = pillow_convert(image, self.rgb_order)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
elif self.backend == 'cv2':
|
||||
image = cv2.imdecode(
|
||||
np.array(bytearray(image_data), dtype='uint8'),
|
||||
cv2.IMREAD_COLOR)
|
||||
if self.rgb_order == 'RGB':
|
||||
try:
|
||||
cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image)
|
||||
except Exception as e:
|
||||
print(img_path, e)
|
||||
else:
|
||||
image = Image.open(img_path).convert(self.rgb_order)
|
||||
image_np = np.asarray(image).transpose(
|
||||
(2, 0, 1)) # Tensor type needs shape to be (C, H, W)
|
||||
image = torch.from_numpy(image_np)
|
||||
return image
|
||||
|
||||
def __call__(self, item):
|
||||
for key in self.file_keys:
|
||||
if 'prefix' in item['meta']:
|
||||
img_path = os.path.join(item['meta']['prefix'],
|
||||
item['meta'][key])
|
||||
else:
|
||||
img_path = item['meta'][key]
|
||||
item[key.replace('_path', '')] = self.read_file(img_path)
|
||||
item['meta']['rgb_order'] = self.rgb_order
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
LoadImageFromFileList.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class LoadPILImageFromFile(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
rgb_order = cfg.get('RGB_ORDER', 'RGB')
|
||||
assert rgb_order in ('RGB', 'BGR')
|
||||
self.rgb_order = rgb_order
|
||||
|
||||
def __call__(self, item):
|
||||
if 'prefix' in item['meta']:
|
||||
img_path = os.path.join(item['meta']['prefix'],
|
||||
item['meta']['img_path'])
|
||||
else:
|
||||
img_path = item['meta']['img_path']
|
||||
|
||||
with FS.get_from(img_path) as img_path:
|
||||
image = Image.open(img_path).convert(self.rgb_order)
|
||||
item['img'] = image
|
||||
item['meta']['rgb_order'] = self.rgb_order
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'RGB_ORDER': {
|
||||
'value': 'RGB',
|
||||
'description': 'rgb order'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class LoadCvImageFromFile(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
rgb_order = cfg.get('RGB_ORDER', 'RGB')
|
||||
assert rgb_order in ('RGB', 'BGR')
|
||||
self.rgb_order = rgb_order
|
||||
|
||||
def __call__(self, item):
|
||||
if 'prefix' in item['meta']:
|
||||
img_path = os.path.join(item['meta']['prefix'],
|
||||
item['meta']['img_path'])
|
||||
else:
|
||||
img_path = item['meta']['img_path']
|
||||
|
||||
with FS.get_from(img_path) as img_path:
|
||||
image = cv2.imread(img_path, cv2.IMREAD_COLOR)
|
||||
if self.rgb_order == 'RGB':
|
||||
cv2.cvtColor(image, cv2.COLOR_BGR2RGB, image)
|
||||
item['img'] = image
|
||||
item['meta']['rgb_order'] = self.rgb_order
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'RGB_ORDER': {
|
||||
'value': 'RGB',
|
||||
'description': 'rgb order'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,486 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import numbers
|
||||
import os.path as osp
|
||||
import queue
|
||||
import random
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from scepter.modules.transform import LoadImageFromFile
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.file_system import DATA_FS as FS
|
||||
from scepter.modules.utils.video_reader.frame_sampler import do_frame_sample
|
||||
from scepter.modules.utils.video_reader.video_reader import VideoReaderWrapper
|
||||
|
||||
|
||||
def _interval_based_sampling(vid_length,
|
||||
vid_fps,
|
||||
target_fps,
|
||||
clip_idx,
|
||||
num_clips,
|
||||
num_frames,
|
||||
interval,
|
||||
minus_interval=False):
|
||||
""" Generates the frame index list using interval based sampling.
|
||||
|
||||
Args:
|
||||
vid_length (int): The length of the whole video (valid selection range).
|
||||
vid_fps (float): The original video fps.
|
||||
target_fps (int): The target decode fps.
|
||||
clip_idx (int): -1 for random temporal sampling, and positive values for sampling specific clip from the video.
|
||||
num_clips (int): The total clips to be sampled from each video.
|
||||
Combined with clip_idx, the sampled video is the "clip_idx-th" video from "num_clips" videos.
|
||||
num_frames (int): Number of frames in each sampled clips.
|
||||
interval (int): The interval to sample each frame.
|
||||
minus_interval (bool):
|
||||
|
||||
Returns:
|
||||
index (torch.Tensor): The sampled frame indexes.
|
||||
"""
|
||||
if num_frames == 1:
|
||||
index = [random.randint(0, vid_length - 1)]
|
||||
else:
|
||||
# transform FPS
|
||||
clip_length = num_frames * interval * vid_fps / target_fps
|
||||
|
||||
max_idx = max(vid_length - clip_length, 0)
|
||||
if clip_idx == -1: # random sampling
|
||||
start_idx = random.uniform(0, max_idx)
|
||||
else:
|
||||
if num_clips == 1:
|
||||
start_idx = max_idx / 2
|
||||
else:
|
||||
start_idx = max_idx * clip_idx / num_clips
|
||||
if minus_interval:
|
||||
end_idx = start_idx + clip_length - interval
|
||||
else:
|
||||
end_idx = start_idx + clip_length - 1
|
||||
|
||||
index = torch.linspace(start_idx, end_idx, num_frames)
|
||||
index = torch.clamp(index, 0, vid_length - 1).long()
|
||||
|
||||
return index
|
||||
|
||||
|
||||
def _segment_based_sampling(vid_length, clip_idx, num_clips, num_frames,
|
||||
random_sample):
|
||||
""" Generates the frame index list using segment based sampling.
|
||||
|
||||
Args:
|
||||
vid_length (int): The length of the whole video (valid selection range).
|
||||
clip_idx (int): -1 for random temporal sampling, and positive values for sampling specific clip from the video.
|
||||
num_clips (int): The total clips to be sampled from each video.
|
||||
Combined with clip_idx, the sampled video is the "clip_idx-th" video from "num_clips" videos.
|
||||
num_frames (int): Number of frames in each sampled clips.
|
||||
random_sample (bool): Whether or not to randomly sample from each segment. True for train and False for test.
|
||||
|
||||
Returns:
|
||||
index (torch.Tensor): The sampled frame indexes.
|
||||
"""
|
||||
index = torch.zeros(num_frames)
|
||||
index_range = torch.linspace(0, vid_length, num_frames + 1)
|
||||
for idx in range(num_frames):
|
||||
if random_sample:
|
||||
index[idx] = random.uniform(index_range[idx], index_range[idx + 1])
|
||||
else:
|
||||
if num_clips == 1:
|
||||
index[idx] = (index_range[idx] + index_range[idx + 1]) / 2
|
||||
else:
|
||||
index[idx] = index_range[idx] + (index_range[
|
||||
idx + 1] - index_range[idx]) * (clip_idx + 1) / num_clips
|
||||
index = torch.round(torch.clamp(index, 0, vid_length - 1)).long()
|
||||
|
||||
return index
|
||||
|
||||
|
||||
R = threading.Lock()
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class LoadVideoFromFile(object):
|
||||
""" Open video file, extract frames, convert to tensor.
|
||||
|
||||
Args:
|
||||
num_frames (int): T dimension value.
|
||||
sample_type (str): See
|
||||
`from essmc2.metric.video_reader import FRAME_SAMPLERS; print(FRAME_SAMPLERS)` to get candidates,
|
||||
default is 'interval'.
|
||||
clip_duration (Optional[float]): Needed for 'interval' sampling type.
|
||||
decoder (str): Video decoder name, default is decord.
|
||||
|
||||
"""
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.num_frames = cfg.NUM_FRAMES
|
||||
self.sample_type = cfg.get('SAMPLE_TYPE', 'interval')
|
||||
self.clip_duration = cfg.get('CLIP_DURATION', None)
|
||||
|
||||
assert self.sample_type in ('uniform', 'interval', 'segment'), \
|
||||
f'Expected sample type in (uniform, interval, segment), got {self.sample_type}'
|
||||
if self.sample_type == 'interval':
|
||||
assert isinstance(self.clip_duration, numbers.Number), \
|
||||
'Interval style sampling needs clip_duration not None'
|
||||
self.decoder = cfg.get('DECODER', 'decord')
|
||||
|
||||
def __call__(self, item):
|
||||
"""
|
||||
Args:
|
||||
item (dict):
|
||||
item['meta']['prefix'] (Optional[str]): Prefix of video_path.
|
||||
item['meta']['video_path'] (str): Required.
|
||||
item['meta']['clip_id'] (Optional[int]): Multi-view test needs it, default is 0.
|
||||
item['meta']['num_clips'] (Optional[int]): Multi-view test needs it, default is 1.
|
||||
item['meta']['start_sec'] (Optional[float]): Uniform sampling needs it.
|
||||
item['meta']['end_sec'] (Optional[float]): Uniform sampling needs it.
|
||||
|
||||
Returns:
|
||||
item(dict):
|
||||
item['video'] (torch.Tensor): a THWC tensor.
|
||||
"""
|
||||
meta = item['meta']
|
||||
video_path = meta['video_path'] if 'prefix' not in meta else osp.join(
|
||||
meta['prefix'], meta['video_path'])
|
||||
|
||||
with FS.get_from(video_path) as local_path:
|
||||
vr = VideoReaderWrapper(local_path, decoder=self.decoder)
|
||||
|
||||
params = dict()
|
||||
clip_id = meta.get('clip_id') or 0
|
||||
num_clips = meta.get('num_clips') or 1
|
||||
if self.sample_type == 'interval':
|
||||
# default is test mode for interval and segment
|
||||
params.update(clip_duration=self.clip_duration,
|
||||
clip_id=clip_id,
|
||||
num_clips=num_clips)
|
||||
elif self.sample_type == 'segment':
|
||||
# default is test mode for interval and segment
|
||||
params.update(clip_id=clip_id, num_clips=num_clips)
|
||||
else:
|
||||
# uniform, needs start_sec, clip_duration or end_sec
|
||||
start_sec = meta['start_sec'] - meta['start_sec']
|
||||
if 'end_sec' in meta:
|
||||
end_sec = meta['end_sec'] - meta['start_sec']
|
||||
elif self.clip_duration is not None:
|
||||
end_sec = start_sec + self.clip_duration
|
||||
else:
|
||||
raise ValueError(
|
||||
'Uniform sampling needs start_sec & end_sec / start_sec & clip_duration'
|
||||
)
|
||||
params.update(start_sec=start_sec, end_sec=end_sec)
|
||||
|
||||
decode_list = do_frame_sample(self.sample_type, vr.len, vr.fps,
|
||||
self.num_frames, **params)
|
||||
item['video'] = vr.sample_frames(decode_list)
|
||||
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'SAMPLE_TYPE': {
|
||||
'value': 'interval',
|
||||
'description': 'sample type'
|
||||
},
|
||||
'CLIP_DURATION': {
|
||||
'value': None,
|
||||
'description': 'clip duration'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class LoadVideoFromFrameList(object):
|
||||
""" extract frames, convert to tensor.
|
||||
|
||||
Args:
|
||||
num_frames (int): T dimension value.
|
||||
sample_type (str): See
|
||||
`from essmc2.metric.video_reader import FRAME_SAMPLERS; print(FRAME_SAMPLERS)` to get candidates,
|
||||
default is 'interval'.
|
||||
clip_duration (Optional[float]): Needed for 'interval' sampling type.
|
||||
decoder (str): Video decoder name, default is decord.
|
||||
"""
|
||||
para_dict = [{
|
||||
'NUM_FRAMES': {
|
||||
'value': 30,
|
||||
'description': 'clip length!'
|
||||
},
|
||||
'SAMPLE_TYPE': {
|
||||
'value': 'interval',
|
||||
'description': 'sample type'
|
||||
},
|
||||
'CLIP_DURATION': {
|
||||
'value': None,
|
||||
'description': 'clip duration'
|
||||
},
|
||||
'RGB_ORDER': {
|
||||
'value': 'RGB',
|
||||
'description': 'rgb order'
|
||||
},
|
||||
'BACKEND': {
|
||||
'value': 'pillow',
|
||||
'description': 'input backend'
|
||||
}
|
||||
}]
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.load_ins = LoadImageFromFile(cfg, logger=logger)
|
||||
self.num_frames = cfg.NUM_FRAMES
|
||||
self.sample_type = cfg.get('SAMPLE_TYPE', 'interval')
|
||||
self.clip_duration = cfg.get('CLIP_DURATION', None)
|
||||
|
||||
assert self.sample_type in ('uniform', 'interval', 'segment'), \
|
||||
f'Expected sample type in (uniform, interval, segment), got {self.sample_type}'
|
||||
if self.sample_type == 'interval':
|
||||
assert isinstance(self.clip_duration, numbers.Number), \
|
||||
'Interval style sampling needs clip_duration not None'
|
||||
|
||||
def __call__(self, item):
|
||||
"""
|
||||
Args:
|
||||
item (dict):
|
||||
item['meta']['video_path'] (str): Required.
|
||||
item['meta']['clip_id'] (Optional[int]): Multi-view test needs it, default is 0.
|
||||
item['meta']['num_clips'] (Optional[int]): Multi-view test needs it, default is 1.
|
||||
item['meta']['start_sec'] (Optional[float]): Uniform sampling needs it.
|
||||
item['meta']['end_sec'] (Optional[float]): Uniform sampling needs it.
|
||||
item['meta']['frames'](list): all frames file path list
|
||||
|
||||
Returns:
|
||||
item(dict):
|
||||
item['video'] (torch.Tensor): a THWC tensor.
|
||||
"""
|
||||
meta = item['meta']
|
||||
params = dict()
|
||||
clip_id = meta.get('clip_id') or 0
|
||||
num_clips = meta.get('num_clips') or 1
|
||||
if self.sample_type == 'interval':
|
||||
# default is test mode for interval and segment
|
||||
params.update(clip_duration=self.clip_duration,
|
||||
clip_id=clip_id,
|
||||
num_clips=num_clips)
|
||||
elif self.sample_type == 'segment':
|
||||
# default is test mode for interval and segment
|
||||
params.update(clip_id=clip_id, num_clips=num_clips)
|
||||
else:
|
||||
# uniform, needs start_sec, clip_duration or end_sec
|
||||
start_sec = meta['start_sec'] - meta['start_sec']
|
||||
if 'end_sec' in meta:
|
||||
end_sec = meta['end_sec'] - meta['start_sec']
|
||||
elif self.clip_duration is not None:
|
||||
end_sec = start_sec + self.clip_duration
|
||||
else:
|
||||
raise ValueError(
|
||||
'Uniform sampling needs start_sec & end_sec / start_sec & clip_duration'
|
||||
)
|
||||
params.update(start_sec=start_sec, end_sec=end_sec)
|
||||
frames = meta['frames']
|
||||
fps = meta['fps']
|
||||
decode_list = do_frame_sample(self.sample_type, len(frames), fps,
|
||||
self.num_frames, **params)
|
||||
sample_frames = [{'meta': {'img_path': frame}} for frame in frames]
|
||||
img_path_queue = queue.Queue()
|
||||
[
|
||||
img_path_queue.put_nowait([idx, item])
|
||||
for idx, item in enumerate(sample_frames)
|
||||
]
|
||||
img_queue = queue.Queue()
|
||||
|
||||
def download_file():
|
||||
while not img_path_queue.empty():
|
||||
R.acquire()
|
||||
try:
|
||||
idx, item = img_path_queue.get_nowait()
|
||||
except Exception:
|
||||
R.release()
|
||||
continue
|
||||
R.release()
|
||||
img_queue.put_nowait([idx, self.load_ins(item)])
|
||||
|
||||
threading_list = []
|
||||
for _ in range(8):
|
||||
t = threading.Thread(target=download_file)
|
||||
t.daemon = True
|
||||
t.start()
|
||||
threading_list.append(t)
|
||||
[th.join() for th in threading_list]
|
||||
# print(f"one video download time {time.time() - st}")
|
||||
sample_frames = []
|
||||
while not img_queue.empty():
|
||||
sample_frames.append(img_queue.get_nowait())
|
||||
sample_frames.sort(key=lambda x: x[0])
|
||||
# sample_frames = [self.load_ins(item) for item in sample_frames]
|
||||
item['video'] = np.array(
|
||||
[sample_frames[frame_id][1]['img'] for frame_id in decode_list])
|
||||
# item['video'] = item['video'].transpose([0, 3, 1, 2])
|
||||
item['meta'].pop('frames')
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
LoadVideoFromFrameList.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class DecodeVideoToTensor(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
""" DecodeVideoToTensor
|
||||
Args:
|
||||
num_frames (int): Decode frames number.
|
||||
target_fps (int): Decode frames fps, default is 30.
|
||||
sample_mode (str): Interval or segment sampling, default is interval.
|
||||
sample_interval (int): Sample interval between output frames for interval sample mode, default is 4.
|
||||
sample_minus_interval (bool): If minus interval for interval sample mode, default is False.
|
||||
repeat (int): Number of clips to be decoded from each video, if repeat > 1, outputs will be named like
|
||||
'video-0', 'video-1'. Normally, 1 for classification task, 2 for contrastive learning.
|
||||
"""
|
||||
import decord
|
||||
from decord import VideoReader
|
||||
self.VideoReader = VideoReader
|
||||
decord.bridge.set_bridge('torch')
|
||||
self.num_frames = cfg.NUM_FRAMES
|
||||
self.target_fps = cfg.get('TARGET_FPS', 30)
|
||||
self.sample_mode = cfg.get('SAMPLE_MODE', 'interval')
|
||||
self.sample_interval = cfg.get('SAMPLE_INTERVAL', 4)
|
||||
self.sample_minus_interval = cfg.get('SAMPLE_MINUS_INTERVAL', False)
|
||||
self.repeat = cfg.get('REPEAT', 1)
|
||||
|
||||
def __call__(self, item):
|
||||
""" Call to invoke decode
|
||||
Args:
|
||||
item (dict): A dict contains which file to decode and how to decode.
|
||||
Normally, it has structure like
|
||||
{
|
||||
"meta": {
|
||||
"prefix" (str, None): if not None, prefix will be added to video_path.
|
||||
"video_path" (str): Absolute (prefix is None) or relative path.
|
||||
"clip_idx" (int): -1 means random sampling, >=0 means do temporal crop.
|
||||
"num_clips" (int): if clip_idx >= 0, clip_idx must < num_clips
|
||||
}
|
||||
}
|
||||
|
||||
Returns:
|
||||
A dict contains original input item and "video" tensor.
|
||||
"""
|
||||
meta = item['meta']
|
||||
video_path = meta['video_path'] \
|
||||
if 'prefix' not in meta else osp.join(meta['prefix'], meta['video_path'])
|
||||
|
||||
with FS.get_from(video_path) as local_path:
|
||||
vr = self.VideoReader(local_path)
|
||||
# default is test mode
|
||||
clip_id = meta.get('clip_id') or 0
|
||||
num_clips = meta.get('num_clips') or 1
|
||||
|
||||
vid_len = len(vr)
|
||||
vid_fps = vr.get_avg_fps()
|
||||
|
||||
frame_list = []
|
||||
for _ in range(self.repeat):
|
||||
if self.sample_mode == 'interval':
|
||||
decode_list = _interval_based_sampling(
|
||||
vid_len, vid_fps, self.target_fps, clip_id, num_clips,
|
||||
self.num_frames, self.sample_interval,
|
||||
self.sample_minus_interval)
|
||||
else:
|
||||
decode_list = _segment_based_sampling(
|
||||
vid_len, clip_id, num_clips, self.num_frames,
|
||||
clip_id == -1)
|
||||
|
||||
# Decord gives inconsistent result for avi files. Getting full frames will fix it, although slower.
|
||||
# See https://github.com/dmlc/decord/issues/195
|
||||
if video_path.lower().endswith('avi'):
|
||||
full_decode_list = list(
|
||||
range(0,
|
||||
torch.max(decode_list).item() + 1))
|
||||
full_frames = vr.get_batch(full_decode_list)
|
||||
frames = full_frames[decode_list].clone()
|
||||
else:
|
||||
frames = vr.get_batch(decode_list).clone()
|
||||
|
||||
frame_list.append(frames)
|
||||
|
||||
if self.repeat == 1:
|
||||
item['video'] = frame_list[0]
|
||||
else:
|
||||
for idx, frame_tensor in zip(range(self.repeat), frame_list):
|
||||
item[f'video-{idx}'] = frame_tensor
|
||||
|
||||
del vr
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'NUM_FRAMES': {
|
||||
'value': 30,
|
||||
'description': 'num frame'
|
||||
},
|
||||
'TARGET_FPS': {
|
||||
'value': 30,
|
||||
'description': 'target fps'
|
||||
},
|
||||
'SAMPLE_MODE': {
|
||||
'value': 'interval',
|
||||
'description': 'sample mode'
|
||||
},
|
||||
'SAMPLE_INTERVAL': {
|
||||
'value': 4,
|
||||
'description': 'sample interval'
|
||||
},
|
||||
'SAMPLE_MINUS_INTERVAL': {
|
||||
'value': False,
|
||||
'description': 'sample minus interval'
|
||||
},
|
||||
'REPEAT': {
|
||||
'value': 1,
|
||||
'description': 'repeat'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,50 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
from scepter.modules.utils.config import Config
|
||||
from scepter.modules.utils.registry import Registry, build_from_config
|
||||
|
||||
|
||||
def build_pipeline(pipeline, registry, logger=None, *args, **kwargs):
|
||||
|
||||
if isinstance(pipeline, list):
|
||||
if len(pipeline) == 0:
|
||||
return build_from_config(Config(cfg_dict={'NAME': 'Identity'},
|
||||
load=False),
|
||||
registry,
|
||||
logger=logger,
|
||||
*args,
|
||||
**kwargs)
|
||||
elif len(pipeline) == 1:
|
||||
return build_pipeline(pipeline[0], registry, logger, *args,
|
||||
**kwargs)
|
||||
else:
|
||||
return build_from_config(Config(cfg_dict={
|
||||
'NAME': 'Compose',
|
||||
'TRANSFORMS': pipeline
|
||||
},
|
||||
load=False),
|
||||
registry,
|
||||
logger=logger,
|
||||
*args,
|
||||
**kwargs)
|
||||
elif isinstance(pipeline, Config):
|
||||
return build_from_config(pipeline,
|
||||
registry,
|
||||
logger=logger,
|
||||
*args,
|
||||
**kwargs)
|
||||
elif pipeline is None:
|
||||
return build_from_config(Config(cfg_dict={'NAME': 'Identity'},
|
||||
load=False),
|
||||
registry,
|
||||
logger=logger,
|
||||
*args,
|
||||
**kwargs)
|
||||
else:
|
||||
raise TypeError(
|
||||
f'Expect pipeline_cfg to be dict or list or None, got {type(pipeline)}'
|
||||
)
|
||||
|
||||
|
||||
TRANSFORMS = Registry('TRANSFORMS', build_func=build_pipeline)
|
||||
@@ -0,0 +1,198 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.distribute import we
|
||||
|
||||
|
||||
def to_tensor(data):
|
||||
if isinstance(data, torch.Tensor):
|
||||
return data
|
||||
elif isinstance(data, np.ndarray):
|
||||
return torch.from_numpy(data)
|
||||
elif isinstance(data, list):
|
||||
return torch.tensor(data)
|
||||
elif isinstance(data, int):
|
||||
return torch.LongTensor([data])
|
||||
elif isinstance(data, float):
|
||||
return torch.FloatTensor([data])
|
||||
else:
|
||||
raise TypeError(f'Unsupported type {type(data)}')
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class ToTensor(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.keys = cfg.KEYS
|
||||
|
||||
def __call__(self, item):
|
||||
for key in self.keys:
|
||||
item[key] = to_tensor(item[key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{'KEYS': {'value': [], 'description': 'keys'}}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class Select(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.keys = cfg.KEYS
|
||||
meta_keys = cfg.get('META_KEYS', [])
|
||||
if not isinstance(meta_keys, (list, tuple)):
|
||||
raise TypeError(
|
||||
f'Expected meta_keys to be list or tuple, got {type(meta_keys)}'
|
||||
)
|
||||
self.meta_keys = meta_keys
|
||||
|
||||
def __call__(self, item):
|
||||
data = {}
|
||||
for key in self.keys:
|
||||
data[key] = item[key]
|
||||
if 'meta' in item and len(self.meta_keys) > 0:
|
||||
data['meta'] = {}
|
||||
for key in self.meta_keys:
|
||||
data['meta'][key] = item['meta'][key]
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'KEYS': {
|
||||
'value': [],
|
||||
'description': 'keys'
|
||||
},
|
||||
'META_KEYS': {
|
||||
'value': [],
|
||||
'description': 'meta keys'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class Rename(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.in_keys = cfg.IN_KEYS
|
||||
self.out_keys = cfg.OUT_KEYS
|
||||
|
||||
def __call__(self, item):
|
||||
data = {}
|
||||
for idx, key in enumerate(self.in_keys):
|
||||
data[self.out_keys[idx]] = item[key]
|
||||
have_key_set = set(self.in_keys)
|
||||
for k, v in item.items():
|
||||
if k not in have_key_set:
|
||||
data[k] = v
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'IN_KEYS': {
|
||||
'value': [],
|
||||
'description':
|
||||
'The keys need to rename, the other keys are outputed by default.'
|
||||
},
|
||||
'OUT_KEYS': {
|
||||
'value': [],
|
||||
'description':
|
||||
'The keys need to rename, the other keys are outputed by default.'
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class TensorToGPU(object):
|
||||
def __init__(self, cfg, logger=None):
|
||||
self.keys = cfg.KEYS
|
||||
self.device_id = we.rank
|
||||
|
||||
def __call__(self, item):
|
||||
ret = {}
|
||||
for key, value in item.items():
|
||||
if key in self.keys and isinstance(
|
||||
value, torch.Tensor) and torch.cuda.is_available():
|
||||
ret[key] = value.cuda(self.device_id, non_blocking=True)
|
||||
else:
|
||||
ret[key] = value
|
||||
return ret
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
para_dict = [{
|
||||
'KEYS': {
|
||||
'value': [],
|
||||
'description': 'keys'
|
||||
},
|
||||
'DEVICE_ID': {
|
||||
'value':
|
||||
0,
|
||||
'description':
|
||||
"device id, which should be set according to current GPU's rank"
|
||||
}
|
||||
}]
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,84 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import numbers
|
||||
|
||||
import numpy as np
|
||||
import opencv_transforms.functional as cv2_TF
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from PIL.Image import Image
|
||||
|
||||
from scepter.modules.transform import TRANSFORMS, ImageTransform
|
||||
from scepter.modules.transform.image import BACKENDS
|
||||
from scepter.modules.transform.utils import BACKEND_PILLOW, BACKEND_TORCHVISION
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class FlexibleCropXL(ImageTransform):
|
||||
para_dict = [{
|
||||
'IS_CENTER': {
|
||||
'value': False,
|
||||
'description': 'Use center crop or not.'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(ImageTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(FlexibleCropXL, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
self.size = cfg.get('SIZE', None)
|
||||
self.is_center = cfg.get('IS_CENTER', False)
|
||||
if self.size is not None:
|
||||
if isinstance(self.size, numbers.Number):
|
||||
self.size = [self.size, self.size]
|
||||
self.callable = TF.crop if self.backend in (
|
||||
BACKEND_PILLOW, BACKEND_TORCHVISION) else cv2_TF.crop
|
||||
|
||||
def __call__(self, item):
|
||||
if isinstance(self.input_key, str):
|
||||
self.input_key = [self.input_key]
|
||||
if isinstance(self.output_key, str):
|
||||
self.output_key = [self.output_key]
|
||||
meta = item.get('meta', {})
|
||||
for idx, key in enumerate(self.input_key):
|
||||
self.check_image_type(item[key])
|
||||
if isinstance(item[key], (torch.Tensor, np.ndarray)):
|
||||
if self.backend in (BACKEND_PILLOW, BACKEND_TORCHVISION):
|
||||
w, h, c = item[key].shape
|
||||
else:
|
||||
h, w, c = item[key].shape
|
||||
elif isinstance(item[key], Image):
|
||||
w, h = item[key].size
|
||||
if 'image_size' in meta:
|
||||
oh, ow = meta['image_size']
|
||||
else:
|
||||
assert self.size is not None
|
||||
oh, ow = self.size
|
||||
meta['image_size'] = [oh, ow]
|
||||
|
||||
delta_h = h - oh
|
||||
delta_w = w - ow
|
||||
if not self.is_center:
|
||||
top = np.random.randint(0, delta_h + 1)
|
||||
left = np.random.randint(0, delta_w + 1)
|
||||
else:
|
||||
top = delta_h // 2
|
||||
left = delta_w // 2
|
||||
|
||||
item[self.output_key[idx]] = self.callable(item[key], top, left,
|
||||
oh, ow)
|
||||
item[self.output_key[idx] + '_' +
|
||||
'original_size_as_tuple'] = torch.tensor([h, w])
|
||||
item[self.output_key[idx] + '_' +
|
||||
'target_size_as_tuple'] = torch.tensor([oh, ow])
|
||||
item[self.output_key[idx] + '_' +
|
||||
'crop_coords_top_left'] = torch.tensor([top, left])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
FlexibleCropXL.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,66 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from packaging import version
|
||||
from PIL import Image
|
||||
from torchvision.version import __version__ as tv_version
|
||||
|
||||
try:
|
||||
import accimage
|
||||
except ImportError:
|
||||
accimage = None
|
||||
|
||||
|
||||
def is_pil_image(img):
|
||||
if accimage is not None:
|
||||
return isinstance(img, (Image.Image, accimage.Image))
|
||||
else:
|
||||
return isinstance(img, Image.Image)
|
||||
|
||||
|
||||
def is_cv2_image(img):
|
||||
return isinstance(img, np.ndarray) and img.dtype == np.uint8
|
||||
|
||||
|
||||
def is_tensor(t):
|
||||
return isinstance(t, torch.Tensor)
|
||||
|
||||
|
||||
INPUT_PIL_TYPE_WARNING = 'input should be PIL Image'
|
||||
INPUT_CV2_TYPE_WARNING = 'input should be cv2 image(uint8 np.ndarray)'
|
||||
INPUT_TENSOR_TYPE_WARNING = 'input should be tensor(uint8 np.ndarray)'
|
||||
|
||||
# Recommend to use nn.Module backend to transform
|
||||
TORCHVISION_CAPABILITY = version.parse(tv_version) >= version.parse('0.8.0')
|
||||
|
||||
BACKEND_TORCHVISION = 'torchvision'
|
||||
BACKEND_PILLOW = 'pillow'
|
||||
BACKEND_CV2 = 'cv2'
|
||||
|
||||
# Recommend to use InterpolationMode since torchvision 0.9.0
|
||||
INTERPOLATION_MODE_CAPABILITY = version.parse(tv_version) >= version.parse(
|
||||
'0.9.0')
|
||||
if INTERPOLATION_MODE_CAPABILITY:
|
||||
from torchvision.transforms.functional import InterpolationMode
|
||||
else:
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings('ignore', message='Default upsampling behavior.*')
|
||||
INTERPOLATION_STYLE = {
|
||||
'bilinear':
|
||||
Image.BILINEAR
|
||||
if not INTERPOLATION_MODE_CAPABILITY else InterpolationMode('bilinear'),
|
||||
'nearest':
|
||||
Image.NEAREST
|
||||
if not INTERPOLATION_MODE_CAPABILITY else InterpolationMode('nearest'),
|
||||
'bicubic':
|
||||
Image.BICUBIC
|
||||
if not INTERPOLATION_MODE_CAPABILITY else InterpolationMode('bicubic'),
|
||||
}
|
||||
INTERPOLATION_STYLE_CV2 = {
|
||||
'bilinear': cv2.INTER_LINEAR,
|
||||
'nearest': cv2.INTER_NEAREST,
|
||||
'bicubic': cv2.INTER_CUBIC,
|
||||
}
|
||||
@@ -0,0 +1,561 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision.transforms.functional as functional
|
||||
import torchvision.transforms.transforms as transforms
|
||||
from packaging import version
|
||||
from torchvision.version import __version__ as tv_version
|
||||
|
||||
from scepter.modules.transform.registry import TRANSFORMS
|
||||
from scepter.modules.transform.utils import (BACKEND_TORCHVISION,
|
||||
INTERPOLATION_STYLE, is_tensor)
|
||||
# torchvision.transform._transforms_video is deprecated since torchvision 0.10.0, use transform instead
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
|
||||
use_video_transforms = version.parse(tv_version) < version.parse('0.10.0')
|
||||
|
||||
BACKENDS = (BACKEND_TORCHVISION, )
|
||||
|
||||
|
||||
class VideoTransform(object):
|
||||
para_dict = [{
|
||||
'INPUT_KEY': {
|
||||
'value': 'img',
|
||||
'description': 'input key'
|
||||
},
|
||||
'OUTPUT_KEY': {
|
||||
'value': 'img',
|
||||
'description': 'input key'
|
||||
},
|
||||
'BACKEND': {
|
||||
'value': 'pillow',
|
||||
'description': 'backend'
|
||||
}
|
||||
}]
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
backend = cfg.get('BACKEND', BACKEND_TORCHVISION)
|
||||
self.input_key = cfg.get('INPUT_KEY', 'video')
|
||||
self.output_key = cfg.get('OUTPUT_KEY', 'video')
|
||||
self.backend = backend
|
||||
|
||||
def check_video_type(self, input_video):
|
||||
if self.backend == BACKEND_TORCHVISION:
|
||||
assert is_tensor(input_video)
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
VideoTransform.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class RandomResizedCropVideo(VideoTransform):
|
||||
"""Crop a random portion of video and resize it to a given size.
|
||||
|
||||
Expect the video is a torch tensor with shape [..., H, W]
|
||||
|
||||
Args:
|
||||
size (int or sequence): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the output size will be matched to (size, size).
|
||||
scale (tuple of float): Specifies the lower and upper bounds for the random area of the crop,
|
||||
before resizing. The scale is defined with respect to the area of the original image.
|
||||
ratio (tuple of float): lower and upper bounds for the random aspect ratio of the crop, before
|
||||
resizing.
|
||||
interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported.
|
||||
"""
|
||||
para_dict = [{
|
||||
'SIZE': {
|
||||
'value': 0,
|
||||
'description': 'size'
|
||||
},
|
||||
'SCALE': {
|
||||
'value': [0.08, 1.0],
|
||||
'description': 'scale'
|
||||
},
|
||||
'RATIO': {
|
||||
'value': [3. / 4., 4. / 3.],
|
||||
'description': 'ratio'
|
||||
},
|
||||
'INTERPOLATION': {
|
||||
'value': 'bilinear',
|
||||
'description': 'interpolation'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(VideoTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
size, scale = cfg.SIZE, cfg.get('SCALE', [0.08, 1.0])
|
||||
ratio = cfg.get('RATIO', [3. / 4., 4. / 3.])
|
||||
interpolation = cfg.get('INTERPOLATION', 'bilinear')
|
||||
super(RandomResizedCropVideo, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
self.interpolation = interpolation
|
||||
if isinstance(size, (tuple, list)):
|
||||
assert len(size) == 2
|
||||
size = tuple(size)
|
||||
elif isinstance(size, int):
|
||||
size = (size, size)
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Unexpected type {type(size)}, expected int or tuple or list')
|
||||
|
||||
if use_video_transforms:
|
||||
from torchvision.transforms._transforms_video import \
|
||||
RandomResizedCropVideo as RandomResizedCropVideoOp
|
||||
self.callable = RandomResizedCropVideoOp(size, scale, ratio,
|
||||
self.interpolation)
|
||||
else:
|
||||
self.callable = transforms.RandomResizedCrop(
|
||||
size, scale, ratio, INTERPOLATION_STYLE[self.interpolation])
|
||||
|
||||
def __call__(self, item):
|
||||
self.check_video_type(item[self.input_key])
|
||||
item[self.output_key] = self.callable(item[self.input_key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
RandomResizedCropVideo.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class CenterCropVideo(VideoTransform):
|
||||
""" Crops the given video at the center.
|
||||
|
||||
Expect the video is a torch tensor with shape [..., H, W]
|
||||
|
||||
Args:
|
||||
size (sequence or int): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the output size will be matched to (size, size).
|
||||
"""
|
||||
para_dict = [{'SIZE': {'value': 0, 'description': 'size'}}]
|
||||
para_dict[0].update(VideoTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
size = cfg.SIZE
|
||||
super(CenterCropVideo, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
self.size = size
|
||||
|
||||
if use_video_transforms:
|
||||
from torchvision.transforms._transforms_video import \
|
||||
CenterCropVideo as CenterCropVideoOp
|
||||
self.callable = CenterCropVideoOp(size)
|
||||
else:
|
||||
self.callable = transforms.CenterCrop(size)
|
||||
|
||||
def __call__(self, item):
|
||||
self.check_video_type(item[self.input_key])
|
||||
item[self.output_key] = self.callable(item[self.input_key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
CenterCropVideo.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class RandomHorizontalFlipVideo(VideoTransform):
|
||||
""" Horizontally flip the given video randomly with a given probability.
|
||||
|
||||
Expect the video is a torch tensor with shape [..., H, W]
|
||||
|
||||
Args:
|
||||
p (float): probability of the image being flipped. Default value is 0.5
|
||||
"""
|
||||
para_dict = [{'P': {'value': 0.5, 'description': 'P'}}]
|
||||
para_dict[0].update(VideoTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
p = cfg.get('P', 0.5)
|
||||
super(RandomHorizontalFlipVideo, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
|
||||
if use_video_transforms:
|
||||
from torchvision.transforms._transforms_video import \
|
||||
RandomHorizontalFlipVideo as RandomHorizontalFlipVideoOp
|
||||
self.callable = RandomHorizontalFlipVideoOp(p)
|
||||
else:
|
||||
self.callable = transforms.RandomHorizontalFlip(p)
|
||||
|
||||
def __call__(self, item):
|
||||
self.check_video_type(item[self.input_key])
|
||||
item[self.output_key] = self.callable(item[self.input_key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
RandomHorizontalFlipVideo.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class NormalizeVideo(VideoTransform):
|
||||
""" Normalize a tensor video with mean and standard deviation.
|
||||
Expect the video is a torch tensor with shape [..., H, W]
|
||||
|
||||
Args:
|
||||
mean (sequence): Sequence of means for each channel.
|
||||
std (sequence): Sequence of standard deviations for each channel.
|
||||
"""
|
||||
para_dict = [{
|
||||
'MEAN': {
|
||||
'value': [],
|
||||
'description': 'mean'
|
||||
},
|
||||
'STD': {
|
||||
'value': [],
|
||||
'description': 'std'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(VideoTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
mean, std = cfg.MEAN, cfg.STD
|
||||
super(NormalizeVideo, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
self.mean = np.array(mean, dtype=np.float32)
|
||||
self.std = np.array(std, dtype=np.float32)
|
||||
|
||||
if use_video_transforms:
|
||||
from torchvision.transforms._transforms_video import \
|
||||
NormalizeVideo as NormalizeVideoOp
|
||||
self.callable = NormalizeVideoOp(self.mean, self.std)
|
||||
else:
|
||||
self.callable = transforms.Normalize(self.mean, self.std)
|
||||
|
||||
def __call__(self, item):
|
||||
video = item[self.input_key]
|
||||
if not use_video_transforms:
|
||||
video = video.permute(1, 0, 2, 3)
|
||||
video = self.callable(video)
|
||||
if not use_video_transforms:
|
||||
video = video.permute(1, 0, 2, 3)
|
||||
item[self.output_key] = video
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
NormalizeVideo.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class VideoToTensor(VideoTransform):
|
||||
""" Convert a uint8 type tensor to a float32 tensor, permute it and scale output to [0.0, 1.0].
|
||||
|
||||
Expect the video is a uint8 torch tensor with shape [T, H, W, C]
|
||||
"""
|
||||
para_dict = [{}]
|
||||
para_dict[0].update(VideoTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
super(VideoToTensor, self).__init__(cfg, logger=logger)
|
||||
assert self.backend in BACKENDS
|
||||
|
||||
def __call__(self, item):
|
||||
video = item[self.input_key]
|
||||
if isinstance(video, np.ndarray):
|
||||
video = torch.tensor(video)
|
||||
|
||||
if not torch.is_tensor(video):
|
||||
raise TypeError('video should be Tensor. Got %s' % type(video))
|
||||
|
||||
if not video.ndimension() == 4:
|
||||
raise ValueError('video should be 4D. Got %dD' % video.dim())
|
||||
|
||||
if not video.dtype == torch.uint8:
|
||||
raise TypeError(
|
||||
'video tensor should have data type uint8. Got %s' %
|
||||
str(video.dtype))
|
||||
|
||||
item[self.output_key] = video.float().permute(3, 0, 1, 2) / 255.0
|
||||
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
VideoToTensor.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class AutoResizedCropVideo(VideoTransform):
|
||||
""" Crop the video with a given position and resize it to a given size.
|
||||
|
||||
Expect the video is a torch tensor with shape [..., H, W].
|
||||
|
||||
Input ``crop_mode`` supports values:
|
||||
- `cc`: center-center
|
||||
- `cl`: left-center
|
||||
- `cr`: right-center
|
||||
- `tl`: left-top
|
||||
- `tr`: right-top
|
||||
- `bl`: left-bottom
|
||||
- `br`: right-bottom
|
||||
|
||||
Args:
|
||||
size (int or sequence): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the output size will be matched to (size, size).
|
||||
scale (tuple of float): Specifies the lower and upper bounds for the random area of the crop,
|
||||
before resizing. The scale is defined with respect to the area of the original image.
|
||||
interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported.
|
||||
"""
|
||||
para_dict = [{'SCALE': {'value': [0.08, 1.0], 'description': 'scale'}}]
|
||||
para_dict[0].update(VideoTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
size, scale = cfg.SIZE, cfg.get('SCALE', [0.08, 1.0])
|
||||
interpolation = cfg.get('INTERPOLATION', 'bilinear')
|
||||
super(AutoResizedCropVideo, self).__init__(cfg, logger=logger)
|
||||
if isinstance(size, (tuple, list)):
|
||||
assert len(size) == 2
|
||||
size = tuple(size)
|
||||
elif isinstance(size, int):
|
||||
size = (size, size)
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Unexpected type {type(size)}, expected int or tuple or list')
|
||||
self.size = size
|
||||
self.scale = scale
|
||||
self.interpolation_mode = interpolation
|
||||
|
||||
def get_crop(self, clip, crop_mode='cc'):
|
||||
scale = random.uniform(*self.scale)
|
||||
_, _, video_height, video_width = clip.shape
|
||||
min_length = min(video_height, video_width)
|
||||
crop_size = int(min_length * scale)
|
||||
center_x = video_width // 2
|
||||
center_y = video_height // 2
|
||||
box_half = crop_size // 2
|
||||
|
||||
# default is cc
|
||||
x0 = center_x - box_half
|
||||
y0 = center_y - box_half
|
||||
if crop_mode == 'cl':
|
||||
x0 = 0
|
||||
y0 = center_y - box_half
|
||||
elif crop_mode == 'cr':
|
||||
x0 = video_width - crop_size
|
||||
y0 = center_y - box_half
|
||||
elif crop_mode == 'tl':
|
||||
x0 = 0
|
||||
y0 = 0
|
||||
elif crop_mode == 'tr':
|
||||
x0 = video_width - crop_size
|
||||
y0 = 0
|
||||
elif crop_mode == 'bl':
|
||||
x0 = 0
|
||||
y0 = video_height - crop_size
|
||||
elif crop_mode == 'br':
|
||||
x0 = video_width - crop_size
|
||||
y0 = video_height - crop_size
|
||||
|
||||
if use_video_transforms:
|
||||
from torchvision.transforms.functional import resized_crop
|
||||
return resized_crop(clip, y0, x0, crop_size, crop_size, self.size,
|
||||
self.interpolation_mode)
|
||||
else:
|
||||
return functional.resized_crop(
|
||||
clip, y0, x0, crop_size, crop_size, self.size,
|
||||
INTERPOLATION_STYLE[self.interpolation_mode])
|
||||
|
||||
def __call__(self, item):
|
||||
self.check_video_type(item[self.input_key])
|
||||
crop_mode = item['meta'].get('crop_mode') or 'cc'
|
||||
item[self.output_key] = self.get_crop(item[self.input_key], crop_mode)
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
AutoResizedCropVideo.para_dict,
|
||||
set_name=True)
|
||||
|
||||
|
||||
@TRANSFORMS.register_class()
|
||||
class ResizeVideo(VideoTransform):
|
||||
"""Resize video to a given size.
|
||||
|
||||
Expect the video is a torch tensor with shape [..., H, W].
|
||||
|
||||
Args:
|
||||
size (int or sequence): Desired output size.
|
||||
If size is a sequence like (h, w), the output size will be matched to this.
|
||||
If size is an int, the smaller edge of the image will be matched to this
|
||||
number maintaining the aspect ratio.
|
||||
interpolation (str): Desired interpolation string, 'bilinear', 'nearest', 'bicubic' are supported.
|
||||
"""
|
||||
para_dict = [{
|
||||
'SCALE': {
|
||||
'value': [0.08, 1.0],
|
||||
'description': 'scale'
|
||||
},
|
||||
'INTERPOLATION': {
|
||||
'value': 'bilinear',
|
||||
'description': 'interpolation'
|
||||
}
|
||||
}]
|
||||
para_dict[0].update(VideoTransform.para_dict[0])
|
||||
|
||||
def __init__(self, cfg, logger=None):
|
||||
size = cfg.SIZE
|
||||
interpolation = cfg.get('INTERPOLATION', 'bilinear')
|
||||
super(ResizeVideo, self).__init__(cfg, logger=logger)
|
||||
self.size = size
|
||||
if isinstance(self.size, (tuple, list)):
|
||||
self.size = tuple(self.size)
|
||||
assert len(self.size) == 2
|
||||
else:
|
||||
if not isinstance(self.size, int):
|
||||
raise ValueError(
|
||||
f'Expected size to be tuple or list or int, got {type(self.size)}'
|
||||
)
|
||||
self.interpolation_mode = interpolation
|
||||
|
||||
def resize(self, clip):
|
||||
if use_video_transforms:
|
||||
from torchvision.transforms.functional import resize
|
||||
|
||||
# resize function only takes a tuple size
|
||||
# so we need to compute scaled target size here
|
||||
if isinstance(self.size, int):
|
||||
h, w = clip.shape[-2], clip.shape[-1]
|
||||
if (w <= h and w == self.size) or (h <= w and h == self.size):
|
||||
return clip
|
||||
if w < h:
|
||||
ow = self.size
|
||||
oh = int(self.size * h / w)
|
||||
else:
|
||||
oh = self.size
|
||||
ow = int(self.size * w / h)
|
||||
size = (oh, ow)
|
||||
else:
|
||||
size = self.size
|
||||
return resize(clip, size, self.interpolation_mode)
|
||||
else:
|
||||
return functional.resize(
|
||||
clip, self.size, INTERPOLATION_STYLE[self.interpolation_mode])
|
||||
|
||||
def __call__(self, item):
|
||||
self.check_video_type(item[self.input_key])
|
||||
item[self.output_key] = self.resize(item[self.input_key])
|
||||
return item
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
'''
|
||||
{ "ENV" :
|
||||
{ "description" : "",
|
||||
"A" : {
|
||||
"value": 1.0,
|
||||
"description": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
:return:
|
||||
'''
|
||||
return dict_to_yaml('TRANSFORM',
|
||||
__class__.__name__,
|
||||
ResizeVideo.para_dict,
|
||||
set_name=True)
|
||||
Reference in New Issue
Block a user