PyTorch: Dropping part of gradient dependence

Viewed 25

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?

0 Answers
Related