def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0, padding_mode='zeros',
dilation=1, groups=1, bias=True, device=None,
shared_keys=True, key_mem_units=2, psi_fn='reduce2d', key_size=None, **kwargs):
self.in_channels = in_channels
self.out_channels = out_channels
self.kernel_size = kernel_size if isinstance(kernel_size, Iterable) else (kernel_size, kernel_size)
self.stride = stride if isinstance(stride, Iterable) else (stride, stride)
self.padding = padding
self.padding_mode = padding_mode
self.dilation = dilation if isinstance(dilation, Iterable) else (dilation, dilation)
self.groups = groups
self.bias = bias
self.in_features = math.prod(self.kernel_size) * self.in_channels
valid_padding_modes = {'zeros', 'reflect', 'replicate', 'circular'}
if padding_mode not in valid_padding_modes:
raise ValueError("padding_mode must be one of {}, but got padding_mode='{}'".format(valid_padding_modes,
padding_mode))
if isinstance(padding, str):
self.__reversed_padding_repeated_twice = [0, 0] * len(self.kernel_size)
if padding == 'same':
for d, k, i in zip(self.dilation, self.kernel_size,
range(len(self.kernel_size) - 1, -1, -1)):
total_padding = d * (k - 1)
left_pad = total_padding // 2
self.__reversed_padding_repeated_twice[2 * i] = left_pad
self.__reversed_padding_repeated_twice[2 * i + 1] = (total_padding - left_pad)
else:
self.padding = padding if isinstance(padding, Iterable) else (padding, padding)
self.__reversed_padding_repeated_twice = tuple(x for x in reversed(self.padding) for _ in range(2))
if kwargs is not None:
assert 'q' not in kwargs, "The number of CNUs is automatically determined, do not set argument 'q'"
assert 'd' not in kwargs, "The size of each key can be specified with argument 'key_size', " \
"do not set argument 'd'"
assert 'm' not in kwargs, "The number of keys and memory units can be specified with argument " \
"'key_mem_units', do not set argument 'm'"
assert 'u' not in kwargs, "Size of each memory unit is automatically determined, do not set argument 'u'"
# Number of keys/memory units
kwargs['m'] = key_mem_units
# Size of each key
if key_size is not None:
if isinstance(key_size, (tuple, list)):
key_size = math.prod(key_size)
kwargs['d'] = key_size
else:
kwargs['d'] = (5 * 5 * self.in_channels)
# Function used to compare input against keys
kwargs['psi_fn'] = psi_fn
if not shared_keys:
# Each neuron is an independent cnu, with its own keys and its own memory units
kwargs['q'] = self.out_channels
kwargs['u'] = self.in_features + (1 if self.bias else 0)
else:
# All the CNUs of the layer share the same keys, thus their memory units are concatenated
kwargs['q'] = 1
kwargs['u'] = self.out_channels * (self.in_features + (1 if self.bias else 0))
# Creating neurons
super(Conv2d, self).__init__(**kwargs)
# Switching device
if device is not None:
self.to(device)