| 8 | |
| 9 | |
| 10 | class DiagonalGaussianDistribution(object): |
| 11 | def __init__(self, parameters, deterministic=False): |
| 12 | self.parameters = parameters |
| 13 | self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) |
| 14 | self.logvar = torch.clamp(self.logvar, -30.0, 20.0) |
| 15 | self.deterministic = deterministic |
| 16 | self.std = torch.exp(0.5 * self.logvar) |
| 17 | self.var = torch.exp(self.logvar) |
| 18 | if self.deterministic: |
| 19 | self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device) |
| 20 | |
| 21 | def sample(self): |
| 22 | # x = self.mean + self.std * torch.randn(self.mean.shape).to( |
| 23 | # device=self.parameters.device |
| 24 | # ) |
| 25 | x = self.mean + self.std * torch.randn_like(self.mean) |
| 26 | return x |
| 27 | |
| 28 | def kl(self, other=None): |
| 29 | if self.deterministic: |
| 30 | return torch.Tensor([0.0]) |
| 31 | else: |
| 32 | if other is None: |
| 33 | return 0.5 * torch.sum( |
| 34 | torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, |
| 35 | dim=[1, 2, 3], |
| 36 | ) |
| 37 | else: |
| 38 | return 0.5 * torch.sum( |
| 39 | torch.pow(self.mean - other.mean, 2) / other.var |
| 40 | + self.var / other.var |
| 41 | - 1.0 |
| 42 | - self.logvar |
| 43 | + other.logvar, |
| 44 | dim=[1, 2, 3], |
| 45 | ) |
| 46 | |
| 47 | def nll(self, sample, dims=[1, 2, 3]): |
| 48 | if self.deterministic: |
| 49 | return torch.Tensor([0.0]) |
| 50 | logtwopi = np.log(2.0 * np.pi) |
| 51 | return 0.5 * torch.sum( |
| 52 | logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var, |
| 53 | dim=dims, |
| 54 | ) |
| 55 | |
| 56 | def mode(self): |
| 57 | return self.mean |
| 58 | |
| 59 | |
| 60 | class AbstractRegularizer(nn.Module): |