GRADE / src /Baselines /radarcam-depth /utils /net_utils.py
Bin-0815's picture
Release all GRADE models, checkpoints, and reviewed evaluation code (part 2)
f348660 verified
Raw History Blame Contribute Delete
19.7 kB
import torch
def activation_func(activation_fn):
'''
Select activation function
Arg(s):
activation_fn : str
name of activation function
'''
if 'linear' in activation_fn:
return None
elif 'leaky_relu' in activation_fn:
return torch.nn.LeakyReLU(negative_slope=0.20, inplace=True)
elif 'relu' in activation_fn:
return torch.nn.ReLU()
elif 'elu' in activation_fn:
return torch.nn.ELU()
elif 'sigmoid' in activation_fn:
return torch.nn.Sigmoid()
else:
raise ValueError('Unsupported activation function: {}'.format(activation_fn))
'''
Network layers
'''
class Conv2d(torch.nn.Module):
'''
2D convolution class
Arg(s):
in_channels : int
number of input channels
out_channels : int
number of output channels
kernel_size : int
size of kernel
stride : int
stride of convolution
weight_initializer : str
kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
activation_func : func
activation function after convolution
use_batch_norm : bool
if set, then applied batch normalization
'''
def __init__(self,
in_channels,
out_channels,
kernel_size=3,
stride=1,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
use_batch_norm=False):
super(Conv2d, self).__init__()
self.use_batch_norm = use_batch_norm
padding = kernel_size // 2
self.conv = torch.nn.Conv2d(
in_channels,
out_channels,
kernel_size=kernel_size,
stride=stride,
padding=padding,
bias=False)
# Select the type of weight initialization, by default kaiming_uniform
if weight_initializer == 'kaiming_normal':
torch.nn.init.kaiming_normal_(self.conv.weight)
elif weight_initializer == 'xavier_normal':
torch.nn.init.xavier_normal_(self.conv.weight)
elif weight_initializer == 'xavier_uniform':
torch.nn.init.xavier_uniform_(self.conv.weight)
self.activation_func = activation_func
if self.use_batch_norm:
self.batch_norm = torch.nn.BatchNorm2d(out_channels)
def forward(self, x):
conv = self.conv(x)
conv = self.batch_norm(conv) if self.use_batch_norm else conv
if self.activation_func is not None:
return self.activation_func(conv)
else:
return conv
class TransposeConv2d(torch.nn.Module):
'''
Transpose convolution class
Arg(s):
in_channels : int
number of input channels
out_channels : int
number of output channels
kernel_size : int
size of kernel (k x k)
weight_initializer : str
kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
activation_func : func
activation function after convolution
use_batch_norm : bool
if set, then applied batch normalization
'''
def __init__(self,
in_channels,
out_channels,
kernel_size=3,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
use_batch_norm=False):
super(TransposeConv2d, self).__init__()
self.use_batch_norm = use_batch_norm
padding = kernel_size // 2
self.deconv = torch.nn.ConvTranspose2d(
in_channels,
out_channels,
kernel_size=kernel_size,
stride=2,
padding=padding,
output_padding=1,
bias=False)
# Select the type of weight initialization, by default kaiming_uniform
if weight_initializer == 'kaiming_normal':
torch.nn.init.kaiming_normal_(self.conv.weight)
elif weight_initializer == 'xavier_normal':
torch.nn.init.xavier_normal_(self.conv.weight)
elif weight_initializer == 'xavier_uniform':
torch.nn.init.xavier_uniform_(self.conv.weight)
self.activation_func = activation_func
if self.use_batch_norm:
self.batch_norm = torch.nn.BatchNorm2d(out_channels)
def forward(self, x):
deconv = self.deconv(x)
deconv = self.batch_norm(deconv) if self.use_batch_norm else deconv
if self.activation_func is not None:
return self.activation_func(deconv)
else:
return deconv
class UpConv2d(torch.nn.Module):
'''
Up-convolution (upsample + convolution) block class
Arg(s):
in_channels : int
number of input channels
out_channels : int
number of output channels
shape : list[int]
two element tuple of ints (height, width)
kernel_size : int
size of kernel (k x k)
weight_initializer : str
kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
activation_func : func
activation function after convolution
use_batch_norm : bool
if set, then applied batch normalization
'''
def __init__(self,
in_channels,
out_channels,
kernel_size=3,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
use_batch_norm=False):
super(UpConv2d, self).__init__()
self.conv = Conv2d(
in_channels,
out_channels,
kernel_size=kernel_size,
stride=1,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
def forward(self, x, shape):
upsample = torch.nn.functional.interpolate(x, size=shape)
conv = self.conv(upsample)
return conv
class FullyConnected(torch.nn.Module):
'''
Fully connected layer
Arg(s):
in_channels : int
number of input neurons
out_channels : int
number of output neurons
dropout_rate : float
probability to use dropout
'''
def __init__(self,
in_features,
out_features,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
dropout_rate=0.00):
super(FullyConnected, self).__init__()
self.fully_connected = torch.nn.Linear(in_features, out_features)
if weight_initializer == 'kaiming_normal':
torch.nn.init.kaiming_normal_(self.fully_connected.weight)
elif weight_initializer == 'xavier_normal':
torch.nn.init.xavier_normal_(self.fully_connected.weight)
elif weight_initializer == 'xavier_uniform':
torch.nn.init.xavier_uniform_(self.fully_connected.weight)
self.activation_func = activation_func
if dropout_rate > 0.00 and dropout_rate <= 1.00:
self.dropout = torch.nn.Dropout(p=dropout_rate)
else:
self.dropout = None
def forward(self, x):
fully_connected = self.fully_connected(x)
if self.activation_func is not None:
fully_connected = self.activation_func(fully_connected)
if self.dropout is not None:
return self.dropout(fully_connected)
else:
return fully_connected
'''
Network encoder blocks
'''
class ResNetBlock(torch.nn.Module):
'''
Basic ResNet block class
Arg(s):
in_channels : int
number of input channels
out_channels : int
number of output channels
stride : int
stride of convolution
weight_initializer : str
kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
activation_func : func
activation function after convolution
use_batch_norm : bool
if set, then applied batch normalization
'''
def __init__(self,
in_channels,
out_channels,
stride=1,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
use_batch_norm=False):
super(ResNetBlock, self).__init__()
self.activation_func = activation_func
self.conv1 = Conv2d(
in_channels,
out_channels,
kernel_size=3,
stride=stride,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
self.conv2 = Conv2d(
out_channels,
out_channels,
kernel_size=3,
stride=1,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
self.projection = Conv2d(
in_channels,
out_channels,
kernel_size=1,
stride=stride,
weight_initializer=weight_initializer,
activation_func=None,
use_batch_norm=False)
def forward(self, x):
# Perform 2 convolutions
conv1 = self.conv1(x)
conv2 = self.conv2(conv1)
# Perform projection if (1) shape does not match (2) channels do not match
in_shape = list(x.shape)
out_shape = list(conv2.shape)
if in_shape[2:4] != out_shape[2:4] or in_shape[1] != out_shape[1]:
X = self.projection(x)
else:
X = x
# f(x) + x
return self.activation_func(conv2 + X)
class ResNetBottleneckBlock(torch.nn.Module):
'''
ResNet bottleneck block class
Arg(s):
in_channels : int
number of input channels
out_channels : int
number of output channels
stride : int
stride of convolution
weight_initializer : str
kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
activation_func : func
activation function after convolution
use_batch_norm : bool
if set, then applied batch normalization
'''
def __init__(self,
in_channels,
out_channels,
stride=1,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
use_batch_norm=False):
super(ResNetBottleneckBlock, self).__init__()
self.activation_func = activation_func
self.conv1 = Conv2d(
in_channels,
out_channels,
kernel_size=1,
stride=1,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
self.conv2 = Conv2d(
out_channels,
out_channels,
kernel_size=3,
stride=stride,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
self.conv3 = Conv2d(
out_channels,
4 * out_channels,
kernel_size=1,
stride=1,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
self.projection = Conv2d(
in_channels,
4 * out_channels,
kernel_size=1,
stride=stride,
weight_initializer=weight_initializer,
activation_func=None,
use_batch_norm=False)
def forward(self, x):
# Perform 2 convolutions
conv1 = self.conv1(x)
conv2 = self.conv2(conv1)
conv3 = self.conv3(conv2)
# Perform projection if (1) shape does not match (2) channels do not match
in_shape = list(x.shape)
out_shape = list(conv2.shape)
if in_shape[2:4] != out_shape[2:4] or in_shape[1] != out_shape[1]:
X = self.projection(x)
else:
X = x
# f(x) + x
return self.activation_func(conv3 + X)
class VGGNetBlock(torch.nn.Module):
'''
VGGNet block class
Arg(s):
in_channels : int
number of input channels
out_channels : int
number of output channels
n_conv : int
number of convolution layers
stride : int
stride of convolution
weight_initializer : str
kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
activation_func : func
activation function after convolution
use_batch_norm : bool
if set, then applied batch normalization
'''
def __init__(self,
in_channels,
out_channels,
n_conv=1,
stride=1,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
use_batch_norm=False):
super(VGGNetBlock, self).__init__()
layers = []
for n in range(n_conv - 1):
conv = Conv2d(
in_channels,
out_channels,
kernel_size=3,
stride=1,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
layers.append(conv)
in_channels = out_channels
conv = Conv2d(
in_channels,
out_channels,
kernel_size=3,
stride=stride,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
layers.append(conv)
self.conv_block = torch.nn.Sequential(*layers)
def forward(self, x):
return self.conv_block(x)
'''
Network decoder blocks
'''
class DecoderBlock(torch.nn.Module):
'''
Decoder block with skip connection
Arg(s):
in_channels : int
number of input channels
skip_channels : int
number of skip connection channels
out_channels : int
number of output channels
weight_initializer : str
kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
activation_func : func
activation function after convolution
use_batch_norm : bool
if set, then applied batch normalization
deconv_type : str
deconvolution types: transpose, up
'''
def __init__(self,
in_channels,
skip_channels,
out_channels,
weight_initializer='kaiming_uniform',
activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True),
use_batch_norm=False,
deconv_type='up'):
super(DecoderBlock, self).__init__()
self.skip_channels = skip_channels
self.deconv_type = deconv_type
if deconv_type == 'transpose':
self.deconv = TransposeConv2d(
in_channels,
out_channels,
kernel_size=3,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
elif deconv_type == 'up':
self.deconv = UpConv2d(
in_channels,
out_channels,
kernel_size=3,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
concat_channels = skip_channels + out_channels
self.conv = Conv2d(
concat_channels,
out_channels,
kernel_size=3,
stride=1,
weight_initializer=weight_initializer,
activation_func=activation_func,
use_batch_norm=use_batch_norm)
def forward(self, x, skip=None, shape=None):
'''
Forward input x through a decoder block and fuse with skip connection
Arg(s):
x : torch.Tensor[float32]
N x C x h x w input tensor
skip : torch.Tensor[float32]
N x F x h x w skip connection
shape : tuple[int]
height, width (H, W) tuple denoting output shape
Returns:
torch.Tensor[float32] : N x K x H x W output tensor
'''
if self.deconv_type == 'transpose':
deconv = self.deconv(x)
elif self.deconv_type == 'up':
if skip is not None:
shape = skip.shape[2:4]
elif shape is not None:
pass
else:
n_height, n_width = x.shape[2:4]
shape = (int(2 * n_height), int(2 * n_width))
deconv = self.deconv(x, shape=shape)
if self.skip_channels > 0:
concat = torch.cat([deconv, skip], dim=1)
else:
concat = deconv
return self.conv(concat)
'''
Utility function to pre-process sparse depth and input depth
'''
class OutlierRemoval(object):
'''
Class to perform outlier removal based on depth difference in local neighborhood
Arg(s):
kernel_size : int
local neighborhood to consider
threshold : float
depth difference threshold
'''
def __init__(self, kernel_size=7, threshold=1.5):
self.kernel_size = kernel_size
self.threshold = threshold
def remove_outliers(self, depth):
'''
Removes erroneous measurements from sparse depth
Arg(s):
depth : torch.Tensor[float32]
N x 1 x H x W tensor sparse depth
Returns:
torch.Tensor[float32] : N x 1 x H x W depth
'''
# Get valid locations
validity_map = torch.where(
depth > 0.0,
torch.ones_like(depth),
depth)
# Replace all zeros with large values
max_value = 10 * torch.max(depth)
depth_max_filled = torch.where(
validity_map <= 0,
torch.full_like(depth, fill_value=max_value),
depth)
# For each neighborhood find the smallest value
padding = self.kernel_size // 2
depth_max_filled = torch.nn.functional.pad(
input=depth_max_filled,
pad=(padding, padding, padding, padding),
mode='constant',
value=max_value)
min_values = -torch.nn.functional.max_pool2d(
input=-depth_max_filled,
kernel_size=self.kernel_size,
stride=1,
padding=0)
# If measurement differs a lot from minimum value then remove
validity_map_clean = torch.where(
min_values < depth - self.threshold,
torch.zeros_like(validity_map),
torch.ones_like(validity_map))
# Update depth map
depth_clean = depth * validity_map_clean
return depth_clean