new project from v0.0.1

This commit is contained in:
duanmurs@163.com
2023-12-28 18:23:39 +08:00
parent fda66e7ca6
commit 66f979f7e9
204 changed files with 37941 additions and 0 deletions
+26
View File
@@ -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)
+506
View File
@@ -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)
+39
View File
@@ -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)
+31
View File
@@ -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)
+626
View File
@@ -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)
+341
View File
@@ -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)
+486
View File
@@ -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)
+50
View File
@@ -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)
+198
View File
@@ -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)
+84
View File
@@ -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)
+66
View File
@@ -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,
}
+561
View File
@@ -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)