MCPcopy Create free account
hub / github.com/TencentARC/AnimeGamer / DiagonalGaussianDistribution

Class DiagonalGaussianDistribution

VDM_Decoder/vae_modules/regularizers.py:10–57  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class 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
60class AbstractRegularizer(nn.Module):

Callers 1

forwardMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected