I'm building a probabilistic model q_{phi)(x), which samples points and returns the corresponding likelihood simultaneously. A simple Gaussian example:
class NormalModel(nn.Module):
def __init__(self, mu : torch.Tensor, sigma : torch.Tensor):
super().__init__()
self.dim_ = mu.shape[0]
self.mu_ = nn.Parameter(mu, requires_grad=True)
self.sigma_ = nn.Parameter(sigma, requires_grad=True)
def sample(self, n_sample):
eps = torch.normal(mean=torch.zeros(n_sample, self.dim_), std=torch.ones(n_sample, self.dim_))
sample = self.mu_+ eps*self.sigma_
log_prob = -(0.5 * ((sample - self.mu_)/self.sigma_)**2 + torch.log(self.sigma_) + m.log(2*m.pi)).sum(dim=-1)
log_prob_detach = -(0.5 * ((sample.detach() - self.mu_)/self.sigma_)**2 + torch.log(self.sigma_) + m.log(2*m.pi)).sum(dim=-1)
return sample, log_prob, log_prob_detach
The sample depdends on the model parameters, x_{phi}, and as such the returned log_prob depends on the parameters in two ways, q_{phi}(x_{phi}). I want to write down a loss function that involves both q_{phi}(x_{phi}) and q_{phi}(x), i.e. the likelihood where sample either does or does not depend on the parameters. I know that I can do this by detaching sample before computing the log_prob, as is done in the computation of log_prob_detach above. This however requires performing the calculation for the log_prob twice. Is there some way to accomplish this after the fact, i.e. perform an operation on log_prob to obtain log_prob_detach?