init
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,48 @@
|
||||
from abc import *
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class BaseModel(nn.Module, metaclass=ABCMeta):
|
||||
def __init__(self, last_dim, num_classes=10, simclr_dim=128):
|
||||
super(BaseModel, self).__init__()
|
||||
self.linear = nn.Linear(last_dim, num_classes)
|
||||
self.simclr_layer = nn.Sequential(
|
||||
nn.Linear(last_dim, last_dim),
|
||||
nn.ReLU(),
|
||||
nn.Linear(last_dim, simclr_dim),
|
||||
)
|
||||
self.shift_cls_layer = nn.Linear(last_dim, 2)
|
||||
self.joint_distribution_layer = nn.Linear(last_dim, 4 * num_classes)
|
||||
|
||||
@abstractmethod
|
||||
def penultimate(self, inputs, all_features=False):
|
||||
pass
|
||||
|
||||
def forward(self, inputs, penultimate=False, simclr=False, shift=False, joint=False):
|
||||
_aux = {}
|
||||
_return_aux = False
|
||||
|
||||
features = self.penultimate(inputs)
|
||||
|
||||
output = self.linear(features)
|
||||
|
||||
if penultimate:
|
||||
_return_aux = True
|
||||
_aux['penultimate'] = features
|
||||
|
||||
if simclr:
|
||||
_return_aux = True
|
||||
_aux['simclr'] = self.simclr_layer(features)
|
||||
|
||||
if shift:
|
||||
_return_aux = True
|
||||
_aux['shift'] = self.shift_cls_layer(features)
|
||||
|
||||
if joint:
|
||||
_return_aux = True
|
||||
_aux['joint'] = self.joint_distribution_layer(features)
|
||||
|
||||
if _return_aux:
|
||||
return output, _aux
|
||||
|
||||
return output
|
||||
@@ -0,0 +1,135 @@
|
||||
import torch.nn as nn
|
||||
|
||||
from models.resnet import ResNet18, ResNet34, ResNet50
|
||||
from models.resnet_imagenet import resnet18, resnet50
|
||||
import models.transform_layers as TL
|
||||
from torchvision import transforms
|
||||
|
||||
|
||||
def get_simclr_augmentation(P, image_size):
|
||||
"""
|
||||
Creates positive data for training.
|
||||
|
||||
:param P: parsed arguments
|
||||
:param image_size: size of image
|
||||
:return: transformation
|
||||
"""
|
||||
|
||||
# parameter for resizecrop
|
||||
resize_scale = (P.resize_factor, 1.0) # resize scaling factor
|
||||
if P.resize_fix: # if resize_fix is True, use same scale
|
||||
resize_scale = (P.resize_factor, P.resize_factor)
|
||||
|
||||
# Align augmentation
|
||||
s = P.color_distort
|
||||
color_jitter = TL.ColorJitterLayer(brightness=s*0.8, contrast=s*0.8, saturation=s*0.8, hue=s*0.2, p=0.8)
|
||||
color_gray = TL.RandomColorGrayLayer(p=0.2)
|
||||
resize_crop = TL.RandomResizedCropLayer(scale=resize_scale, size=(image_size[0], image_size[1]))
|
||||
|
||||
#v_flip = transforms.RandomVerticalFlip()
|
||||
#h_flip = transforms.RandomHorizontalFlip()
|
||||
rand_aff = transforms.RandomAffine(degrees=360, translate=(0.2, 0.2))
|
||||
|
||||
# Transform define #
|
||||
if P.dataset == 'imagenet': # Using RandomResizedCrop at PIL transform
|
||||
transform = nn.Sequential(
|
||||
color_jitter,
|
||||
color_gray,
|
||||
)
|
||||
elif P.dataset == 'CNMC':
|
||||
transform = nn.Sequential(
|
||||
color_jitter,
|
||||
color_gray,
|
||||
resize_crop,
|
||||
)
|
||||
else:
|
||||
transform = nn.Sequential(
|
||||
color_jitter,
|
||||
color_gray,
|
||||
resize_crop,
|
||||
)
|
||||
|
||||
return transform
|
||||
|
||||
|
||||
def get_shift_module(P, eval=False):
|
||||
"""
|
||||
Creates shift transformation (negative).
|
||||
|
||||
:param P: parsed arguments
|
||||
:param eval: whether it is an evaluation step or not
|
||||
:return: transformation
|
||||
"""
|
||||
if P.shift_trans_type == 'rotation':
|
||||
shift_transform = TL.Rotation()
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'cutperm':
|
||||
shift_transform = TL.CutPerm()
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'noise':
|
||||
shift_transform = TL.GaussNoise(mean=P.noise_mean, std=P.noise_std)
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'randpers':
|
||||
shift_transform = TL.RandPers(distortion_scale=P.distortion_scale, p=1)
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'sharp':
|
||||
shift_transform = TL.RandomAdjustSharpness(sharpness_factor=P.sharpness_factor, p=1)
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'blur':
|
||||
kernel_size = int(int(P.res.replace('px', ''))*0.1)
|
||||
if kernel_size%2 == 0:
|
||||
kernel_size+=1
|
||||
sigma = (0.1, float(P.blur_sigma))
|
||||
shift_transform = TL.GaussBlur(kernel_size=kernel_size, sigma=sigma)
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'blur_randpers':
|
||||
kernel_size = int(P.res.replace('px', '')) * 0.1
|
||||
sigma = (0.1, float(P.blur_sigma))
|
||||
shift_transform = TL.BlurRandpers(kernel_size=kernel_size, sigma=sigma, distortion_scale=P.distortion_scale, p=1)
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'blur_sharp':
|
||||
kernel_size = int(P.res.replace('px', '')) * 0.1
|
||||
sigma = (0.1, float(P.blur_sigma))
|
||||
shift_transform = TL.BlurSharpness(kernel_size=kernel_size, sigma=sigma, sharpness_factor=P.sharpness_factor, p=1)
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'randpers_sharp':
|
||||
shift_transform = TL.RandpersSharpness(distortion_scale=P.distortion_scale, p=1, sharpness_factor=P.sharpness_factor)
|
||||
K_shift = 4
|
||||
elif P.shift_trans_type == 'blur_randpers_sharp':
|
||||
kernel_size = int(P.res.replace('px', '')) * 0.1
|
||||
sigma = (0.1, float(P.blur_sigma))
|
||||
shift_transform = TL.BlurRandpersSharpness(kernel_size=kernel_size, sigma=sigma, distortion_scale=P.distortion_scale, p=1, sharpness_factor=P.sharpness_factor)
|
||||
K_shift = 4
|
||||
else:
|
||||
shift_transform = nn.Identity()
|
||||
K_shift = 1
|
||||
|
||||
if not eval and not ('sup' in P.mode):
|
||||
assert P.batch_size == int(128/K_shift)
|
||||
|
||||
return shift_transform, K_shift
|
||||
|
||||
|
||||
def get_shift_classifer(model, K_shift):
|
||||
|
||||
model.shift_cls_layer = nn.Linear(model.last_dim, K_shift)
|
||||
|
||||
return model
|
||||
|
||||
|
||||
def get_classifier(mode, n_classes=10):
|
||||
if mode == 'resnet18':
|
||||
classifier = ResNet18(num_classes=n_classes)
|
||||
elif mode == 'resnet34':
|
||||
classifier = ResNet34(num_classes=n_classes)
|
||||
elif mode == 'resnet50':
|
||||
classifier = ResNet50(num_classes=n_classes)
|
||||
elif mode == 'resnet18_imagenet':
|
||||
classifier = resnet18(num_classes=n_classes)
|
||||
elif mode == 'resnet50_imagenet':
|
||||
classifier = resnet50(num_classes=n_classes)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
return classifier
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
'''ResNet in PyTorch.
|
||||
BasicBlock and Bottleneck module is from the original ResNet paper:
|
||||
[1] Kaiming He, Xiangyu Zhang, Shaoqing Ren, Jian Sun
|
||||
Deep Residual Learning for Image Recognition. arXiv:1512.03385
|
||||
PreActBlock and PreActBottleneck module is from the later paper:
|
||||
[2] Kaiming He, Xiangyu Zhang, Shaoqing Ren, Jian Sun
|
||||
Identity Mappings in Deep Residual Networks. arXiv:1603.05027
|
||||
'''
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from models.base_model import BaseModel
|
||||
from models.transform_layers import NormalizeLayer
|
||||
from torch.nn.utils import spectral_norm
|
||||
|
||||
def conv3x3(in_planes, out_planes, stride=1):
|
||||
return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=False)
|
||||
|
||||
|
||||
class BasicBlock(nn.Module):
|
||||
expansion = 1
|
||||
|
||||
def __init__(self, in_planes, planes, stride=1):
|
||||
super(BasicBlock, self).__init__()
|
||||
self.conv1 = conv3x3(in_planes, planes, stride)
|
||||
self.conv2 = conv3x3(planes, planes)
|
||||
self.bn1 = nn.BatchNorm2d(planes)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
|
||||
self.shortcut = nn.Sequential()
|
||||
if stride != 1 or in_planes != self.expansion*planes:
|
||||
self.shortcut = nn.Sequential(
|
||||
nn.Conv2d(in_planes, self.expansion*planes, kernel_size=1, stride=stride, bias=False),
|
||||
nn.BatchNorm2d(self.expansion*planes)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
out = F.relu(self.bn1(self.conv1(x)))
|
||||
out = self.bn2(self.conv2(out))
|
||||
out += self.shortcut(x)
|
||||
out = F.relu(out)
|
||||
return out
|
||||
|
||||
|
||||
class PreActBlock(nn.Module):
|
||||
'''Pre-activation version of the BasicBlock.'''
|
||||
expansion = 1
|
||||
|
||||
def __init__(self, in_planes, planes, stride=1):
|
||||
super(PreActBlock, self).__init__()
|
||||
self.conv1 = conv3x3(in_planes, planes, stride)
|
||||
self.conv2 = conv3x3(planes, planes)
|
||||
self.bn1 = nn.BatchNorm2d(in_planes)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
|
||||
self.shortcut = nn.Sequential()
|
||||
if stride != 1 or in_planes != self.expansion*planes:
|
||||
self.shortcut = nn.Sequential(
|
||||
nn.Conv2d(in_planes, self.expansion*planes, kernel_size=1, stride=stride, bias=False)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
out = F.relu(self.bn1(x))
|
||||
shortcut = self.shortcut(out)
|
||||
out = self.conv1(out)
|
||||
out = self.conv2(F.relu(self.bn2(out)))
|
||||
out += shortcut
|
||||
return out
|
||||
|
||||
|
||||
class Bottleneck(nn.Module):
|
||||
expansion = 4
|
||||
|
||||
def __init__(self, in_planes, planes, stride=1):
|
||||
super(Bottleneck, self).__init__()
|
||||
self.conv1 = nn.Conv2d(in_planes, planes, kernel_size=1, bias=False)
|
||||
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
|
||||
self.conv3 = nn.Conv2d(planes, self.expansion*planes, kernel_size=1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(planes)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
self.bn3 = nn.BatchNorm2d(self.expansion * planes)
|
||||
|
||||
self.shortcut = nn.Sequential()
|
||||
if stride != 1 or in_planes != self.expansion*planes:
|
||||
self.shortcut = nn.Sequential(
|
||||
nn.Conv2d(in_planes, self.expansion*planes, kernel_size=1, stride=stride, bias=False),
|
||||
nn.BatchNorm2d(self.expansion*planes)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
out = F.relu(self.bn1(self.conv1(x)))
|
||||
out = F.relu(self.bn2(self.conv2(out)))
|
||||
out = self.bn3(self.conv3(out))
|
||||
out += self.shortcut(x)
|
||||
out = F.relu(out)
|
||||
return out
|
||||
|
||||
|
||||
class PreActBottleneck(nn.Module):
|
||||
'''Pre-activation version of the original Bottleneck module.'''
|
||||
expansion = 4
|
||||
|
||||
def __init__(self, in_planes, planes, stride=1):
|
||||
super(PreActBottleneck, self).__init__()
|
||||
self.conv1 = nn.Conv2d(in_planes, planes, kernel_size=1, bias=False)
|
||||
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
|
||||
self.conv3 = nn.Conv2d(planes, self.expansion*planes, kernel_size=1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(in_planes)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
self.bn3 = nn.BatchNorm2d(planes)
|
||||
|
||||
self.shortcut = nn.Sequential()
|
||||
if stride != 1 or in_planes != self.expansion*planes:
|
||||
self.shortcut = nn.Sequential(
|
||||
nn.Conv2d(in_planes, self.expansion*planes, kernel_size=1, stride=stride, bias=False)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
out = F.relu(self.bn1(x))
|
||||
shortcut = self.shortcut(out)
|
||||
out = self.conv1(out)
|
||||
out = self.conv2(F.relu(self.bn2(out)))
|
||||
out = self.conv3(F.relu(self.bn3(out)))
|
||||
out += shortcut
|
||||
return out
|
||||
|
||||
|
||||
class ResNet(BaseModel):
|
||||
def __init__(self, block, num_blocks, num_classes=10):
|
||||
last_dim = 512 * block.expansion
|
||||
super(ResNet, self).__init__(last_dim, num_classes)
|
||||
|
||||
self.in_planes = 64
|
||||
self.last_dim = last_dim
|
||||
|
||||
self.normalize = NormalizeLayer()
|
||||
|
||||
self.conv1 = conv3x3(3, 64)
|
||||
self.bn1 = nn.BatchNorm2d(64)
|
||||
|
||||
self.layer1 = self._make_layer(block, 64, num_blocks[0], stride=1)
|
||||
self.layer2 = self._make_layer(block, 128, num_blocks[1], stride=2)
|
||||
self.layer3 = self._make_layer(block, 256, num_blocks[2], stride=2)
|
||||
self.layer4 = self._make_layer(block, 512, num_blocks[3], stride=2)
|
||||
|
||||
def _make_layer(self, block, planes, num_blocks, stride):
|
||||
strides = [stride] + [1]*(num_blocks-1)
|
||||
layers = []
|
||||
for stride in strides:
|
||||
layers.append(block(self.in_planes, planes, stride))
|
||||
self.in_planes = planes * block.expansion
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
def penultimate(self, x, all_features=False):
|
||||
out_list = []
|
||||
|
||||
out = self.normalize(x)
|
||||
out = self.conv1(out)
|
||||
out = self.bn1(out)
|
||||
out = F.relu(out)
|
||||
out_list.append(out)
|
||||
|
||||
out = self.layer1(out)
|
||||
out_list.append(out)
|
||||
out = self.layer2(out)
|
||||
out_list.append(out)
|
||||
out = self.layer3(out)
|
||||
out_list.append(out)
|
||||
out = self.layer4(out)
|
||||
out_list.append(out)
|
||||
|
||||
out = F.avg_pool2d(out, 4)
|
||||
out = out.view(out.size(0), -1)
|
||||
|
||||
if all_features:
|
||||
return out, out_list
|
||||
else:
|
||||
return out
|
||||
|
||||
|
||||
def ResNet18(num_classes):
|
||||
return ResNet(BasicBlock, [2,2,2,2], num_classes=num_classes)
|
||||
|
||||
def ResNet34(num_classes):
|
||||
return ResNet(BasicBlock, [3,4,6,3], num_classes=num_classes)
|
||||
|
||||
def ResNet50(num_classes):
|
||||
return ResNet(Bottleneck, [3,4,6,3], num_classes=num_classes)
|
||||
@@ -0,0 +1,231 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from models.base_model import BaseModel
|
||||
from models.transform_layers import NormalizeLayer
|
||||
|
||||
|
||||
def conv3x3(in_planes, out_planes, stride=1, groups=1, dilation=1):
|
||||
"""3x3 convolution with padding"""
|
||||
return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,
|
||||
padding=dilation, groups=groups, bias=False, dilation=dilation)
|
||||
|
||||
|
||||
def conv1x1(in_planes, out_planes, stride=1):
|
||||
"""1x1 convolution"""
|
||||
return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False)
|
||||
|
||||
|
||||
class BasicBlock(nn.Module):
|
||||
expansion = 1
|
||||
|
||||
def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1,
|
||||
base_width=64, dilation=1, norm_layer=None):
|
||||
super(BasicBlock, self).__init__()
|
||||
if norm_layer is None:
|
||||
norm_layer = nn.BatchNorm2d
|
||||
if groups != 1 or base_width != 64:
|
||||
raise ValueError('BasicBlock only supports groups=1 and base_width=64')
|
||||
if dilation > 1:
|
||||
raise NotImplementedError("Dilation > 1 not supported in BasicBlock")
|
||||
# Both self.conv1 and self.downsample layers downsample the input when stride != 1
|
||||
self.conv1 = conv3x3(inplanes, planes, stride)
|
||||
self.bn1 = norm_layer(planes)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.conv2 = conv3x3(planes, planes)
|
||||
self.bn2 = norm_layer(planes)
|
||||
self.downsample = downsample
|
||||
self.stride = stride
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
|
||||
out = self.conv1(x)
|
||||
out = self.bn1(out)
|
||||
out = self.relu(out)
|
||||
|
||||
out = self.conv2(out)
|
||||
out = self.bn2(out)
|
||||
|
||||
if self.downsample is not None:
|
||||
identity = self.downsample(x)
|
||||
|
||||
out += identity
|
||||
out = self.relu(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class Bottleneck(nn.Module):
|
||||
# Bottleneck in torchvision places the stride for downsampling at 3x3 convolution(self.conv2)
|
||||
# while original implementation places the stride at the first 1x1 convolution(self.conv1)
|
||||
# according to "Deep residual learning for image recognition"https://arxiv.org/abs/1512.03385.
|
||||
# This variant is also known as ResNet V1.5 and improves accuracy according to
|
||||
# https://ngc.nvidia.com/catalog/model-scripts/nvidia:resnet_50_v1_5_for_pytorch.
|
||||
|
||||
expansion = 4
|
||||
|
||||
def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1,
|
||||
base_width=64, dilation=1, norm_layer=None):
|
||||
super(Bottleneck, self).__init__()
|
||||
if norm_layer is None:
|
||||
norm_layer = nn.BatchNorm2d
|
||||
width = int(planes * (base_width / 64.)) * groups
|
||||
# Both self.conv2 and self.downsample layers downsample the input when stride != 1
|
||||
self.conv1 = conv1x1(inplanes, width)
|
||||
self.bn1 = norm_layer(width)
|
||||
self.conv2 = conv3x3(width, width, stride, groups, dilation)
|
||||
self.bn2 = norm_layer(width)
|
||||
self.conv3 = conv1x1(width, planes * self.expansion)
|
||||
self.bn3 = norm_layer(planes * self.expansion)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.downsample = downsample
|
||||
self.stride = stride
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
|
||||
out = self.conv1(x)
|
||||
out = self.bn1(out)
|
||||
out = self.relu(out)
|
||||
|
||||
out = self.conv2(out)
|
||||
out = self.bn2(out)
|
||||
out = self.relu(out)
|
||||
|
||||
out = self.conv3(out)
|
||||
out = self.bn3(out)
|
||||
|
||||
if self.downsample is not None:
|
||||
identity = self.downsample(x)
|
||||
|
||||
out += identity
|
||||
out = self.relu(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class ResNet(BaseModel):
|
||||
def __init__(self, block, layers, num_classes=10,
|
||||
zero_init_residual=False, groups=1, width_per_group=64, replace_stride_with_dilation=None,
|
||||
norm_layer=None):
|
||||
last_dim = 512 * block.expansion
|
||||
super(ResNet, self).__init__(last_dim, num_classes)
|
||||
if norm_layer is None:
|
||||
norm_layer = nn.BatchNorm2d
|
||||
self._norm_layer = norm_layer
|
||||
|
||||
self.inplanes = 64
|
||||
self.dilation = 1
|
||||
if replace_stride_with_dilation is None:
|
||||
# each element in the tuple indicates if we should replace
|
||||
# the 2x2 stride with a dilated convolution instead
|
||||
replace_stride_with_dilation = [False, False, False]
|
||||
if len(replace_stride_with_dilation) != 3:
|
||||
raise ValueError("replace_stride_with_dilation should be None "
|
||||
"or a 3-element tuple, got {}".format(replace_stride_with_dilation))
|
||||
self.groups = groups
|
||||
self.base_width = width_per_group
|
||||
self.conv1 = nn.Conv2d(3, self.inplanes, kernel_size=7, stride=2, padding=3,
|
||||
bias=False)
|
||||
self.bn1 = norm_layer(self.inplanes)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
||||
self.layer1 = self._make_layer(block, 64, layers[0])
|
||||
self.layer2 = self._make_layer(block, 128, layers[1], stride=2,
|
||||
dilate=replace_stride_with_dilation[0])
|
||||
self.layer3 = self._make_layer(block, 256, layers[2], stride=2,
|
||||
dilate=replace_stride_with_dilation[1])
|
||||
self.layer4 = self._make_layer(block, 512, layers[3], stride=2,
|
||||
dilate=replace_stride_with_dilation[2])
|
||||
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
||||
self.normalize = NormalizeLayer()
|
||||
self.last_dim = 512 * block.expansion
|
||||
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
|
||||
elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):
|
||||
nn.init.constant_(m.weight, 1)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
# Zero-initialize the last BN in each residual branch,
|
||||
# so that the residual branch starts with zeros, and each residual block behaves like an identity.
|
||||
# This improves the model by 0.2~0.3% according to https://arxiv.org/abs/1706.02677
|
||||
if zero_init_residual:
|
||||
for m in self.modules():
|
||||
if isinstance(m, Bottleneck):
|
||||
nn.init.constant_(m.bn3.weight, 0)
|
||||
elif isinstance(m, BasicBlock):
|
||||
nn.init.constant_(m.bn2.weight, 0)
|
||||
|
||||
def _make_layer(self, block, planes, blocks, stride=1, dilate=False):
|
||||
norm_layer = self._norm_layer
|
||||
downsample = None
|
||||
previous_dilation = self.dilation
|
||||
if dilate:
|
||||
self.dilation *= stride
|
||||
stride = 1
|
||||
if stride != 1 or self.inplanes != planes * block.expansion:
|
||||
downsample = nn.Sequential(
|
||||
conv1x1(self.inplanes, planes * block.expansion, stride),
|
||||
norm_layer(planes * block.expansion),
|
||||
)
|
||||
|
||||
layers = []
|
||||
layers.append(block(self.inplanes, planes, stride, downsample, self.groups,
|
||||
self.base_width, previous_dilation, norm_layer))
|
||||
self.inplanes = planes * block.expansion
|
||||
for _ in range(1, blocks):
|
||||
layers.append(block(self.inplanes, planes, groups=self.groups,
|
||||
base_width=self.base_width, dilation=self.dilation,
|
||||
norm_layer=norm_layer))
|
||||
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
def penultimate(self, x, all_features=False):
|
||||
# See note [TorchScript super()]
|
||||
out_list = []
|
||||
|
||||
x = self.normalize(x)
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
x = self.maxpool(x)
|
||||
out_list.append(x)
|
||||
|
||||
x = self.layer1(x)
|
||||
out_list.append(x)
|
||||
x = self.layer2(x)
|
||||
out_list.append(x)
|
||||
x = self.layer3(x)
|
||||
out_list.append(x)
|
||||
x = self.layer4(x)
|
||||
out_list.append(x)
|
||||
|
||||
x = self.avgpool(x)
|
||||
x = torch.flatten(x, 1)
|
||||
|
||||
if all_features:
|
||||
return x, out_list
|
||||
else:
|
||||
return x
|
||||
|
||||
|
||||
def _resnet(arch, block, layers, **kwargs):
|
||||
model = ResNet(block, layers, **kwargs)
|
||||
return model
|
||||
|
||||
|
||||
def resnet18(**kwargs):
|
||||
r"""ResNet-18 model from
|
||||
`"Deep Residual Learning for Image Recognition" <https://arxiv.org/pdf/1512.03385.pdf>`_
|
||||
"""
|
||||
return _resnet('resnet18', BasicBlock, [2, 2, 2, 2], **kwargs)
|
||||
|
||||
|
||||
def resnet50(**kwargs):
|
||||
r"""ResNet-50 model from
|
||||
`"Deep Residual Learning for Image Recognition" <https://arxiv.org/pdf/1512.03385.pdf>`_
|
||||
"""
|
||||
return _resnet('resnet50', Bottleneck, [3, 4, 6, 3], **kwargs)
|
||||
@@ -0,0 +1,643 @@
|
||||
import math
|
||||
import numbers
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.autograd import Function
|
||||
from torchvision import transforms
|
||||
|
||||
if torch.__version__ >= '1.4.0':
|
||||
kwargs = {'align_corners': False}
|
||||
else:
|
||||
kwargs = {}
|
||||
|
||||
|
||||
def rgb2hsv(rgb):
|
||||
"""Convert a 4-d RGB tensor to the HSV counterpart.
|
||||
|
||||
Here, we compute hue using atan2() based on the definition in [1],
|
||||
instead of using the common lookup table approach as in [2, 3].
|
||||
Those values agree when the angle is a multiple of 30°,
|
||||
otherwise they may differ at most ~1.2°.
|
||||
|
||||
References
|
||||
[1] https://en.wikipedia.org/wiki/Hue
|
||||
[2] https://www.rapidtables.com/convert/color/rgb-to-hsv.html
|
||||
[3] https://github.com/scikit-image/scikit-image/blob/master/skimage/color/colorconv.py#L212
|
||||
"""
|
||||
|
||||
r, g, b = rgb[:, 0, :, :], rgb[:, 1, :, :], rgb[:, 2, :, :]
|
||||
|
||||
Cmax = rgb.max(1)[0]
|
||||
Cmin = rgb.min(1)[0]
|
||||
delta = Cmax - Cmin
|
||||
|
||||
hue = torch.atan2(math.sqrt(3) * (g - b), 2 * r - g - b)
|
||||
hue = (hue % (2 * math.pi)) / (2 * math.pi)
|
||||
saturate = delta / Cmax
|
||||
value = Cmax
|
||||
hsv = torch.stack([hue, saturate, value], dim=1)
|
||||
hsv[~torch.isfinite(hsv)] = 0.
|
||||
return hsv
|
||||
|
||||
|
||||
def hsv2rgb(hsv):
|
||||
"""Convert a 4-d HSV tensor to the RGB counterpart.
|
||||
|
||||
>>> %timeit hsv2rgb(hsv)
|
||||
2.37 ms ± 13.4 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
|
||||
>>> %timeit rgb2hsv_fast(rgb)
|
||||
298 µs ± 542 ns per loop (mean ± std. dev. of 7 runs, 1000 loops each)
|
||||
>>> torch.allclose(hsv2rgb(hsv), hsv2rgb_fast(hsv), atol=1e-6)
|
||||
True
|
||||
|
||||
References
|
||||
[1] https://en.wikipedia.org/wiki/HSL_and_HSV#HSV_to_RGB_alternative
|
||||
"""
|
||||
h, s, v = hsv[:, [0]], hsv[:, [1]], hsv[:, [2]]
|
||||
c = v * s
|
||||
|
||||
n = hsv.new_tensor([5, 3, 1]).view(3, 1, 1)
|
||||
k = (n + h * 6) % 6
|
||||
t = torch.min(k, 4 - k)
|
||||
t = torch.clamp(t, 0, 1)
|
||||
|
||||
return v - c * t
|
||||
|
||||
|
||||
class RandomResizedCropLayer(nn.Module):
|
||||
def __init__(self, size=None, scale=(0.08, 1.0), ratio=(3. / 4., 4. / 3.)):
|
||||
'''
|
||||
Inception Crop
|
||||
size (tuple): size of fowarding image (C, W, H)
|
||||
scale (tuple): range of size of the origin size cropped
|
||||
ratio (tuple): range of aspect ratio of the origin aspect ratio cropped
|
||||
'''
|
||||
super(RandomResizedCropLayer, self).__init__()
|
||||
|
||||
_eye = torch.eye(2, 3)
|
||||
self.size = size
|
||||
self.register_buffer('_eye', _eye)
|
||||
self.scale = scale
|
||||
self.ratio = ratio
|
||||
|
||||
def forward(self, inputs, whbias=None):
|
||||
_device = inputs.device
|
||||
N = inputs.size(0)
|
||||
_theta = self._eye.repeat(N, 1, 1)
|
||||
|
||||
if whbias is None:
|
||||
whbias = self._sample_latent(inputs)
|
||||
|
||||
_theta[:, 0, 0] = whbias[:, 0]
|
||||
_theta[:, 1, 1] = whbias[:, 1]
|
||||
_theta[:, 0, 2] = whbias[:, 2]
|
||||
_theta[:, 1, 2] = whbias[:, 3]
|
||||
|
||||
grid = F.affine_grid(_theta, inputs.size(), **kwargs).to(_device)
|
||||
output = F.grid_sample(inputs, grid, padding_mode='reflection', **kwargs)
|
||||
if self.size is not None:
|
||||
output = F.adaptive_avg_pool2d(output, self.size)
|
||||
# output = F.adaptive_avg_pool2d(output, self.size)
|
||||
# output = F.adaptive_avg_pool2d(output, (self.size[0], self.size[1]))
|
||||
|
||||
|
||||
return output
|
||||
|
||||
def _clamp(self, whbias):
|
||||
|
||||
w = whbias[:, 0]
|
||||
h = whbias[:, 1]
|
||||
w_bias = whbias[:, 2]
|
||||
h_bias = whbias[:, 3]
|
||||
|
||||
# Clamp with scale
|
||||
w = torch.clamp(w, *self.scale)
|
||||
h = torch.clamp(h, *self.scale)
|
||||
|
||||
# Clamp with ratio
|
||||
w = self.ratio[0] * h + torch.relu(w - self.ratio[0] * h)
|
||||
w = self.ratio[1] * h - torch.relu(self.ratio[1] * h - w)
|
||||
|
||||
# Clamp with bias range: w_bias \in (w - 1, 1 - w), h_bias \in (h - 1, 1 - h)
|
||||
w_bias = w - 1 + torch.relu(w_bias - w + 1)
|
||||
w_bias = 1 - w - torch.relu(1 - w - w_bias)
|
||||
|
||||
h_bias = h - 1 + torch.relu(h_bias - h + 1)
|
||||
h_bias = 1 - h - torch.relu(1 - h - h_bias)
|
||||
|
||||
whbias = torch.stack([w, h, w_bias, h_bias], dim=0).t()
|
||||
|
||||
return whbias
|
||||
|
||||
def _sample_latent(self, inputs):
|
||||
|
||||
_device = inputs.device
|
||||
N, _, width, height = inputs.shape
|
||||
|
||||
# N * 10 trial
|
||||
area = width * height
|
||||
target_area = np.random.uniform(*self.scale, N * 10) * area
|
||||
log_ratio = (math.log(self.ratio[0]), math.log(self.ratio[1]))
|
||||
aspect_ratio = np.exp(np.random.uniform(*log_ratio, N * 10))
|
||||
|
||||
# If doesn't satisfy ratio condition, then do central crop
|
||||
w = np.round(np.sqrt(target_area * aspect_ratio))
|
||||
h = np.round(np.sqrt(target_area / aspect_ratio))
|
||||
cond = (0 < w) * (w <= width) * (0 < h) * (h <= height)
|
||||
w = w[cond]
|
||||
h = h[cond]
|
||||
cond_len = w.shape[0]
|
||||
if cond_len >= N:
|
||||
w = w[:N]
|
||||
h = h[:N]
|
||||
else:
|
||||
w = np.concatenate([w, np.ones(N - cond_len) * width])
|
||||
h = np.concatenate([h, np.ones(N - cond_len) * height])
|
||||
|
||||
w_bias = np.random.randint(w - width, width - w + 1) / width
|
||||
h_bias = np.random.randint(h - height, height - h + 1) / height
|
||||
w = w / width
|
||||
h = h / height
|
||||
|
||||
whbias = np.column_stack([w, h, w_bias, h_bias])
|
||||
whbias = torch.tensor(whbias, device=_device)
|
||||
|
||||
return whbias
|
||||
|
||||
|
||||
class HorizontalFlipRandomCrop(nn.Module):
|
||||
def __init__(self, max_range):
|
||||
super(HorizontalFlipRandomCrop, self).__init__()
|
||||
self.max_range = max_range
|
||||
_eye = torch.eye(2, 3)
|
||||
self.register_buffer('_eye', _eye)
|
||||
|
||||
def forward(self, input, sign=None, bias=None, rotation=None):
|
||||
_device = input.device
|
||||
N = input.size(0)
|
||||
_theta = self._eye.repeat(N, 1, 1)
|
||||
|
||||
if sign is None:
|
||||
sign = torch.bernoulli(torch.ones(N, device=_device) * 0.5) * 2 - 1
|
||||
if bias is None:
|
||||
bias = torch.empty((N, 2), device=_device).uniform_(-self.max_range, self.max_range)
|
||||
_theta[:, 0, 0] = sign
|
||||
_theta[:, :, 2] = bias
|
||||
|
||||
if rotation is not None:
|
||||
_theta[:, 0:2, 0:2] = rotation
|
||||
|
||||
grid = F.affine_grid(_theta, input.size(), **kwargs).to(_device)
|
||||
output = F.grid_sample(input, grid, padding_mode='reflection', **kwargs)
|
||||
|
||||
return output
|
||||
|
||||
def _sample_latent(self, N, device=None):
|
||||
sign = torch.bernoulli(torch.ones(N, device=device) * 0.5) * 2 - 1
|
||||
bias = torch.empty((N, 2), device=device).uniform_(-self.max_range, self.max_range)
|
||||
return sign, bias
|
||||
|
||||
|
||||
class Rotation(nn.Module):
|
||||
def __init__(self, max_range = 4):
|
||||
super(Rotation, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.prob = 0.5
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
_device = input.device
|
||||
|
||||
_, _, H, W = input.size()
|
||||
|
||||
if aug_index is None:
|
||||
aug_index = np.random.randint(4)
|
||||
|
||||
output = torch.rot90(input, aug_index, (2, 3))
|
||||
|
||||
_prob = input.new_full((input.size(0),), self.prob)
|
||||
_mask = torch.bernoulli(_prob).view(-1, 1, 1, 1)
|
||||
output = _mask * input + (1-_mask) * output
|
||||
|
||||
else:
|
||||
aug_index = aug_index % self.max_range
|
||||
output = torch.rot90(input, aug_index, (2, 3))
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class RandomAdjustSharpness(nn.Module):
|
||||
def __init__(self, sharpness_factor=0.5, p=0.5):
|
||||
super(RandomAdjustSharpness, self).__init__()
|
||||
self.sharpness_factor = sharpness_factor
|
||||
self.prob = p
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
_device = input.device
|
||||
|
||||
_, _, H, W = input.size()
|
||||
if aug_index == 0:
|
||||
output = input
|
||||
else:
|
||||
output = transforms.RandomAdjustSharpness(sharpness_factor=self.sharpness_factor, p=self.prob)(input)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class RandPers(nn.Module):
|
||||
def __init__(self, distortion_scale=0.5, p=0.5):
|
||||
super(RandPers, self).__init__()
|
||||
self.distortion_scale = distortion_scale
|
||||
self.prob = p
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
_device = input.device
|
||||
|
||||
_, _, H, W = input.size()
|
||||
if aug_index == 0:
|
||||
output = input
|
||||
else:
|
||||
output = transforms.RandomPerspective(distortion_scale=self.distortion_scale, p=self.prob)(input)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class GaussBlur(nn.Module):
|
||||
def __init__(self, max_range = 4, kernel_size=3, sigma=(0.1, 2.0)):
|
||||
super(GaussBlur, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.prob = 0.5
|
||||
self.sigma = sigma
|
||||
self.kernel_size = kernel_size
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
_device = input.device
|
||||
|
||||
_, _, H, W = input.size()
|
||||
if aug_index is None:
|
||||
aug_index = np.random.randint(4)
|
||||
|
||||
output = transforms.GaussianBlur(kernel_size=13, sigma=abs(aug_index)+1)(input)
|
||||
|
||||
_prob = input.new_full((input.size(0),), self.prob)
|
||||
_mask = torch.bernoulli(_prob).view(-1, 1, 1, 1)
|
||||
output = _mask * input + (1-_mask) * output
|
||||
|
||||
else:
|
||||
if aug_index == 0:
|
||||
output = input
|
||||
else:
|
||||
output = transforms.GaussianBlur(kernel_size=self.kernel_size, sigma=self.sigma)(input)
|
||||
|
||||
return output
|
||||
|
||||
class GaussNoise(nn.Module):
|
||||
def __init__(self, mean = 0, std = 1):
|
||||
super(GaussNoise, self).__init__()
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
_device = input.device
|
||||
|
||||
_, _, H, W = input.size()
|
||||
|
||||
if aug_index == 0:
|
||||
output = input
|
||||
else:
|
||||
output = input + (torch.randn(input.size()) * self.std + self.mean).to(_device)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class BlurRandpers(nn.Module):
|
||||
def __init__(self, max_range=2, kernel_size=3, sigma=(10, 20), distortion_scale=0.6, p=1):
|
||||
super(BlurRandpers, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.sigma = sigma
|
||||
self.kernel_size = kernel_size
|
||||
self.distortion_scale = distortion_scale
|
||||
self.p = p
|
||||
self.gauss = GaussBlur(kernel_size=self.kernel_size, sigma=self.sigma)
|
||||
self.randpers = RandPers(distortion_scale=self.distortion_scale, p=self.p)
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
output = self.gauss.forward(input=input, aug_index=aug_index)
|
||||
output = self.randpers.forward(input=output, aug_index=aug_index)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class BlurSharpness(nn.Module):
|
||||
def __init__(self, max_range=2, kernel_size=3, sigma=(10, 20), sharpness_factor=0.6, p=1):
|
||||
super(BlurSharpness, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.sigma = sigma
|
||||
self.kernel_size = kernel_size
|
||||
self.sharpness_factor = sharpness_factor
|
||||
self.p = p
|
||||
self.gauss = GaussBlur(kernel_size=self.kernel_size, sigma=self.sigma)
|
||||
self.sharp = RandomAdjustSharpness(sharpness_factor=self.sharpness_factor, p=self.p)
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
output = self.gauss.forward(input=input, aug_index=aug_index)
|
||||
output = self.sharp.forward(input=output, aug_index=aug_index)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class RandpersSharpness(nn.Module):
|
||||
def __init__(self, max_range=2, distortion_scale=0.6, p=1, sharpness_factor=0.6):
|
||||
super(RandpersSharpness, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.distortion_scale = distortion_scale
|
||||
self.p = p
|
||||
self.sharpness_factor = sharpness_factor
|
||||
self.randpers = RandPers(distortion_scale=self.distortion_scale, p=self.p)
|
||||
self.sharp = RandomAdjustSharpness(sharpness_factor=self.sharpness_factor, p=self.p)
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
output = self.randpers.forward(input=input, aug_index=aug_index)
|
||||
output = self.sharp.forward(input=output, aug_index=aug_index)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class BlurRandpersSharpness(nn.Module):
|
||||
def __init__(self, max_range=2, kernel_size=3, sigma=(10, 20), distortion_scale=0.6, p=1, sharpness_factor=0.6):
|
||||
super(BlurRandpersSharpness, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.sigma = sigma
|
||||
self.kernel_size = kernel_size
|
||||
self.distortion_scale = distortion_scale
|
||||
self.p = p
|
||||
self.sharpness_factor = sharpness_factor
|
||||
self.gauss = GaussBlur(kernel_size=self.kernel_size, sigma=self.sigma)
|
||||
self.randpers = RandPers(distortion_scale=self.distortion_scale, p=self.p)
|
||||
self.sharp = RandomAdjustSharpness(sharpness_factor=self.sharpness_factor, p=self.p)
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
output = self.gauss.forward(input=input, aug_index=aug_index)
|
||||
output = self.randpers.forward(input=output, aug_index=aug_index)
|
||||
output = self.sharp.forward(input=output, aug_index=aug_index)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class FourCrop(nn.Module):
|
||||
def __init__(self, max_range = 4):
|
||||
super(FourCrop, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.prob = 0.5
|
||||
|
||||
def forward(self, inputs):
|
||||
outputs = inputs
|
||||
for i in range(8):
|
||||
outputs[i] = self._crop(inputs.size(), inputs[i], i)
|
||||
|
||||
return outputs
|
||||
|
||||
def _crop(self, size, input, i):
|
||||
_, _, H, W = size
|
||||
h_mid = int(H / 2)
|
||||
w_mid = int(W / 2)
|
||||
|
||||
if i == 0 or i == 4:
|
||||
corner = input[:, 0:h_mid, 0:w_mid]
|
||||
elif i == 1 or i == 5:
|
||||
corner = input[:, 0:h_mid, w_mid:]
|
||||
elif i == 2 or i == 6:
|
||||
corner = input[:, h_mid:, 0:w_mid]
|
||||
elif i == 3 or i == 7:
|
||||
corner = input[:, h_mid:, w_mid:]
|
||||
else:
|
||||
corner = input
|
||||
corner = transforms.Resize(size=2*h_mid)(corner)
|
||||
|
||||
return corner
|
||||
|
||||
|
||||
class CutPerm(nn.Module):
|
||||
def __init__(self, max_range = 4):
|
||||
super(CutPerm, self).__init__()
|
||||
self.max_range = max_range
|
||||
self.prob = 0.5
|
||||
|
||||
def forward(self, input, aug_index=None):
|
||||
_device = input.device
|
||||
|
||||
_, _, H, W = input.size()
|
||||
|
||||
if aug_index is None:
|
||||
aug_index = np.random.randint(4)
|
||||
|
||||
output = self._cutperm(input, aug_index)
|
||||
|
||||
_prob = input.new_full((input.size(0),), self.prob)
|
||||
_mask = torch.bernoulli(_prob).view(-1, 1, 1, 1)
|
||||
output = _mask * input + (1 - _mask) * output
|
||||
|
||||
else:
|
||||
aug_index = aug_index % self.max_range
|
||||
output = self._cutperm(input, aug_index)
|
||||
|
||||
return output
|
||||
|
||||
def _cutperm(self, inputs, aug_index):
|
||||
|
||||
_, _, H, W = inputs.size()
|
||||
h_mid = int(H / 2)
|
||||
w_mid = int(W / 2)
|
||||
|
||||
jigsaw_h = aug_index // 2
|
||||
jigsaw_v = aug_index % 2
|
||||
|
||||
if jigsaw_h == 1:
|
||||
inputs = torch.cat((inputs[:, :, h_mid:, :], inputs[:, :, 0:h_mid, :]), dim=2)
|
||||
if jigsaw_v == 1:
|
||||
inputs = torch.cat((inputs[:, :, :, w_mid:], inputs[:, :, :, 0:w_mid]), dim=3)
|
||||
|
||||
return inputs
|
||||
|
||||
|
||||
def assemble(a, b, c, d):
|
||||
ab = torch.cat((a, b), dim=2)
|
||||
cd = torch.cat((c, d), dim=2)
|
||||
output = torch.cat((ab, cd), dim=3)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def quarter(inputs):
|
||||
_, _, H, W = inputs.size()
|
||||
h_mid = int(H / 2)
|
||||
w_mid = int(W / 2)
|
||||
quarters = []
|
||||
quarters.append(inputs[:, :, 0:h_mid, 0:w_mid])
|
||||
quarters.append(inputs[:, :, 0:h_mid, w_mid:])
|
||||
quarters.append(inputs[:, :, h_mid:, 0:w_mid])
|
||||
quarters.append(inputs[:, :, h_mid:, w_mid:])
|
||||
|
||||
return quarters
|
||||
|
||||
|
||||
class HorizontalFlipLayer(nn.Module):
|
||||
def __init__(self):
|
||||
"""
|
||||
img_size : (int, int, int)
|
||||
Height and width must be powers of 2. E.g. (32, 32, 1) or
|
||||
(64, 128, 3). Last number indicates number of channels, e.g. 1 for
|
||||
grayscale or 3 for RGB
|
||||
"""
|
||||
super(HorizontalFlipLayer, self).__init__()
|
||||
|
||||
_eye = torch.eye(2, 3)
|
||||
self.register_buffer('_eye', _eye)
|
||||
|
||||
def forward(self, inputs):
|
||||
_device = inputs.device
|
||||
|
||||
N = inputs.size(0)
|
||||
_theta = self._eye.repeat(N, 1, 1)
|
||||
r_sign = torch.bernoulli(torch.ones(N, device=_device) * 0.5) * 2 - 1
|
||||
_theta[:, 0, 0] = r_sign
|
||||
grid = F.affine_grid(_theta, inputs.size(), **kwargs).to(_device)
|
||||
inputs = F.grid_sample(inputs, grid, padding_mode='reflection', **kwargs)
|
||||
|
||||
return inputs
|
||||
|
||||
|
||||
class RandomColorGrayLayer(nn.Module):
|
||||
def __init__(self, p):
|
||||
super(RandomColorGrayLayer, self).__init__()
|
||||
self.prob = p
|
||||
|
||||
_weight = torch.tensor([[0.299, 0.587, 0.114]])
|
||||
self.register_buffer('_weight', _weight.view(1, 3, 1, 1))
|
||||
|
||||
def forward(self, inputs, aug_index=None):
|
||||
|
||||
if aug_index == 0:
|
||||
return inputs
|
||||
|
||||
l = F.conv2d(inputs, self._weight)
|
||||
gray = torch.cat([l, l, l], dim=1)
|
||||
|
||||
if aug_index is None:
|
||||
_prob = inputs.new_full((inputs.size(0),), self.prob)
|
||||
_mask = torch.bernoulli(_prob).view(-1, 1, 1, 1)
|
||||
|
||||
gray = inputs * (1 - _mask) + gray * _mask
|
||||
|
||||
return gray
|
||||
|
||||
|
||||
class ColorJitterLayer(nn.Module):
|
||||
def __init__(self, p, brightness, contrast, saturation, hue):
|
||||
super(ColorJitterLayer, self).__init__()
|
||||
self.prob = p
|
||||
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)
|
||||
|
||||
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 - value, center + value]
|
||||
if clip_first_on_zero:
|
||||
value[0] = max(value[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 adjust_contrast(self, x):
|
||||
if self.contrast:
|
||||
factor = x.new_empty(x.size(0), 1, 1, 1).uniform_(*self.contrast)
|
||||
means = torch.mean(x, dim=[2, 3], keepdim=True)
|
||||
x = (x - means) * factor + means
|
||||
return torch.clamp(x, 0, 1)
|
||||
|
||||
def adjust_hsv(self, x):
|
||||
f_h = x.new_zeros(x.size(0), 1, 1)
|
||||
f_s = x.new_ones(x.size(0), 1, 1)
|
||||
f_v = x.new_ones(x.size(0), 1, 1)
|
||||
|
||||
if self.hue:
|
||||
f_h.uniform_(*self.hue)
|
||||
if self.saturation:
|
||||
f_s = f_s.uniform_(*self.saturation)
|
||||
if self.brightness:
|
||||
f_v = f_v.uniform_(*self.brightness)
|
||||
|
||||
return RandomHSVFunction.apply(x, f_h, f_s, f_v)
|
||||
|
||||
def transform(self, inputs):
|
||||
# Shuffle transform
|
||||
if np.random.rand() > 0.5:
|
||||
transforms = [self.adjust_contrast, self.adjust_hsv]
|
||||
else:
|
||||
transforms = [self.adjust_hsv, self.adjust_contrast]
|
||||
|
||||
for t in transforms:
|
||||
inputs = t(inputs)
|
||||
|
||||
return inputs
|
||||
|
||||
def forward(self, inputs):
|
||||
_prob = inputs.new_full((inputs.size(0),), self.prob)
|
||||
_mask = torch.bernoulli(_prob).view(-1, 1, 1, 1)
|
||||
return inputs * (1 - _mask) + self.transform(inputs) * _mask
|
||||
|
||||
|
||||
class RandomHSVFunction(Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x, f_h, f_s, f_v):
|
||||
# ctx is a context object that can be used to stash information
|
||||
# for backward computation
|
||||
x = rgb2hsv(x)
|
||||
h = x[:, 0, :, :]
|
||||
h += (f_h * 255. / 360.)
|
||||
h = (h % 1)
|
||||
x[:, 0, :, :] = h
|
||||
x[:, 1, :, :] = x[:, 1, :, :] * f_s
|
||||
x[:, 2, :, :] = x[:, 2, :, :] * f_v
|
||||
x = torch.clamp(x, 0, 1)
|
||||
x = hsv2rgb(x)
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
# We return as many input gradients as there were arguments.
|
||||
# Gradients of non-Tensor arguments to forward must be None.
|
||||
grad_input = None
|
||||
if ctx.needs_input_grad[0]:
|
||||
grad_input = grad_output.clone()
|
||||
return grad_input, None, None, None
|
||||
|
||||
|
||||
class NormalizeLayer(nn.Module):
|
||||
"""
|
||||
In order to certify radii in original coordinates rather than standardized coordinates, we
|
||||
add the Gaussian noise _before_ standardizing, which is why we have standardization be the first
|
||||
layer of the classifier rather than as a part of preprocessing as is typical.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super(NormalizeLayer, self).__init__()
|
||||
|
||||
def forward(self, inputs):
|
||||
return (inputs - 0.5) / 0.5
|
||||
|
||||
Reference in New Issue
Block a user