0% found this document useful (0 votes)
4 views3 pages

Bayesian Linear and Conv2D Layers

Uploaded by

harveyxiangnus
Copyright
© All Rights Reserved
We take content rights seriously. If you suspect this is your content, claim it here.
Available Formats
Download as TXT, PDF, TXT or read online on Scribd
0% found this document useful (0 votes)
4 views3 pages

Bayesian Linear and Conv2D Layers

Uploaded by

harveyxiangnus
Copyright
© All Rights Reserved
We take content rights seriously. If you suspect this is your content, claim it here.
Available Formats
Download as TXT, PDF, TXT or read online on Scribd

"""

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

You might also like