KL divergence between normal distributions parameterized by mean and log-variance.
(mean1, logvar1, mean2, logvar2)
| 21 | models |
| 22 | ''' |
| 23 | def normal_kl(mean1, logvar1, mean2, logvar2): |
| 24 | """ |
| 25 | KL divergence between normal distributions parameterized by mean and log-variance. |
| 26 | """ |
| 27 | return 0.5 * (-1.0 + logvar2 - logvar1 + torch.exp(logvar1 - logvar2) |
| 28 | + (mean1 - mean2)**2 * torch.exp(-logvar2)) |
| 29 | |
| 30 | def discretized_gaussian_log_likelihood(x, *, means, log_scales): |
| 31 | # Assumes data is integers [0, 1] |
nothing calls this directly
no outgoing calls
no test coverage detected