"""
class BBB_LRT_Linear([Link]):
"""
# Bayesian Linear layer with Local Reparameterization Trick.
"""
def __init__(self, in_features, out_features, bias=True, priors=None):
super(BBB_LRT_Linear, self).__init__()
self.in_features = in_features
self.out_features = out_features
self.use_bias = bias
# Initialize weight and bias parameters
self.weight_mu = [Link]([Link](out_features, in_features))
self.weight_log_sigma = [Link]([Link](out_features,
in_features))
if self.use_bias:
self.bias_mu = [Link]([Link](out_features))
self.bias_log_sigma = [Link]([Link](out_features))
else:
self.register_parameter('bias_mu', None)
self.register_parameter('bias_log_sigma', None)
# Initialize prior distributions
if priors is None:
priors = {
'prior_mu': 0,
'prior_sigma': 0.1,
'posterior_mu_initial': (0, 0.1),
'posterior_rho_initial': (-3, 0.1),
}
self.prior_mu = priors['prior_mu']
self.prior_sigma = priors['prior_sigma']
self.posterior_mu_initial = priors['posterior_mu_initial']
self.posterior_rho_initial = priors['posterior_rho_initial']
self.reset_parameters()
def reset_parameters(self):
# Initialize weight and bias means and log standard deviations
[Link].normal_(self.weight_mu, self.posterior_mu_initial[0],
self.posterior_mu_initial[1])
if self.use_bias:
[Link].normal_(self.bias_mu, self.posterior_mu_initial[0],
self.posterior_mu_initial[1])
[Link].constant_(self.weight_log_sigma,
[Link]([Link](self.posterior_rho_initial[0]) - 1))
if self.use_bias:
[Link].constant_(self.bias_log_sigma,
[Link]([Link](self.posterior_rho_initial[0]) - 1))
def forward(self, input, sample=False, calculate_log_probs=False):
if [Link] or sample:
# Calculate mean and variance of activations
act_mu = [Link](input, self.weight_mu, self.bias_mu)
act_var = [Link](input**2, [Link](self.weight_log_sigma)**2, None)
if self.use_bias:
act_var += [Link](self.bias_log_sigma)**2
act_std = [Link](act_var)
# Sample from the activation distribution
eps = torch.randn_like(act_std)
output = act_mu + act_std * eps
else:
output = [Link](input, self.weight_mu, self.bias_mu)
if calculate_log_probs:
log_prior = utils.log_gaussian(self.weight_mu, self.prior_mu,
self.prior_sigma).sum()
log_variational_posterior = utils.log_gaussian(self.weight_mu,
self.weight_mu, [Link](self.weight_log_sigma)).sum()
if self.use_bias:
log_prior += utils.log_gaussian(self.bias_mu, self.prior_mu,
self.prior_sigma).sum()
log_variational_posterior += utils.log_gaussian(self.bias_mu,
self.bias_mu, [Link](self.bias_log_sigma)).sum()
return output, log_prior, log_variational_posterior
else:
return output
class BBB_LRT_Conv2d([Link]):
"""
# Bayesian Conv2d layer with Local Reparameterization Trick.
"""
def __init__(self, in_channels, out_channels, kernel_size,
stride=1, padding=0, dilation=1, groups=1,
bias=True, padding_mode='zeros', priors=None):
super(BBB_LRT_Conv2d, self).__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.kernel_size = (kernel_size, kernel_size)
[Link] = stride
[Link] = padding
[Link] = dilation
[Link] = groups
self.use_bias = bias
self.padding_mode = padding_mode
# Initialize weight and bias parameters
self.weight_mu = [Link]([Link](out_channels, in_channels //
groups, *self.kernel_size))
self.weight_log_sigma = [Link]([Link](out_channels, in_channels
// groups, *self.kernel_size))
if self.use_bias:
self.bias_mu = [Link]([Link](out_channels))
self.bias_log_sigma = [Link]([Link](out_channels))
else:
self.register_parameter('bias_mu', None)
self.register_parameter('bias_log_sigma', None)
# Initialize prior distributions
if priors is None:
priors = {
'prior_mu': 0,
'prior_sigma': 0.1,
'posterior_mu_initial': (0, 0.1),
'posterior_rho_initial': (-3, 0.1),
}
self.prior_mu = priors['prior_mu']
self.prior_sigma = priors['prior_sigma']
self.posterior_mu_initial = priors['posterior_mu_initial']
self.posterior_rho_initial = priors['posterior_rho_initial']
self.reset_parameters()
def reset_parameters(self):
# Initialize weight and bias means and log standard deviations
[Link].normal_(self.weight_mu, self.posterior_mu_initial[0],
self.posterior_mu_initial[1])
if self.use_bias:
[Link].normal_(self.bias_mu, self.posterior_mu_initial[0],
self.posterior_mu_initial[1])
[Link].constant_(self.weight_log_sigma,
[Link]([Link](self.posterior_rho_initial[0]) - 1))
if self.use_bias:
[Link].constant_(self.bias_log_sigma,
[Link]([Link](self.posterior_rho_initial[0]) - 1))
def forward(self, input, sample=False, calculate_log_probs=False):
if [Link] or sample:
# Calculate mean and variance of activations
act_mu = F.conv2d(input, self.weight_mu, self.bias_mu, [Link],
[Link], [Link], [Link])
act_var = F.conv2d(input**2, [Link](self.weight_log_sigma)**2, None,
[Link],
[Link], [Link], [Link])
if self.use_bias:
act_var += [Link](self.bias_log_sigma)[None, :, None, None]**2
act_std = [Link](act_var)
# Sample from the activation distribution
eps = torch.randn_like(act_std)
output = act_mu + act_std * eps
else:
output = F.conv2d(input, self.weight_mu, self.bias_mu, [Link],
[Link], [Link], [Link])
if calculate_log_probs:
log_prior = utils.log_gaussian(self.weight_mu, self.prior_mu,
self.prior_sigma).sum()
log_variational_posterior = utils.log_gaussian(self.weight_mu,
self.weight_mu, [Link](self.weight_log_sigma)).sum()
if self.use_bias:
log_prior += utils.log_gaussian(self.bias_mu, self.prior_mu,
self.prior_sigma).sum()
log_variational_posterior += utils.log_gaussian(self.bias_mu,
self.bias_mu, [Link](self.bias_log_sigma)).sum()
return output, log_prior, log_variational_posterior
else:
return output