MCPcopy Create free account
hub / github.com/DSL-Lab/StreamSplat / Truncated_Gaussian_Model

Class Truncated_Gaussian_Model

model/mixture_model_utils.py:9–132  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7eps = 1e-3
8
9class Truncated_Gaussian_Model(nn.Module):
10 def __init__(self, n_sample=1, nr_mix=1):
11 super(Truncated_Gaussian_Model, self).__init__()
12 self.n_sample = n_sample
13 self.nr_mix = nr_mix
14 self.log_scales_min = -5.0
15 self.log_scales_max = 2.0
16 self.log_scales_a = (self.log_scales_min + self.log_scales_max) / 2
17 self.log_scales_b = (self.log_scales_max - self.log_scales_min) / 2
18 self.perform_sampling = True
19
20 def activate_mean(self, means, mean_activation='tanh'):
21 if mean_activation == 'tanh':
22 return torch.tanh(means) * (1.0 - eps) # (-1, 1)
23 else:
24 raise ValueError(f"Unknown activation function: {self.mean_activation}")
25
26 def expand_params(self, logits, means, log_scales, mean_activation='tanh'):
27 """
28 Expand the parameters to n_samples.
29 """
30 B, N, _ = means.shape # [B, N, dim*nr_mix]
31 dim = int(means.shape[-1] / self.nr_mix)
32 logits = logits.repeat(1, 1, self.n_sample).reshape(B, -1, 1, self.nr_mix) # [B, N*n_sample, 1, nr_mix]
33
34 means = means.reshape(B, -1, dim, self.nr_mix) # [B, N, dim, nr_mix]
35 log_scales = log_scales.reshape(B, -1, dim, self.nr_mix) # [B, N, dim, nr_mix]
36
37 means = means.repeat(1, 1, self.n_sample, 1).reshape(means.shape[0], -1, dim, self.nr_mix) # [B, N*n_sample, dim, nr_mix]
38 log_scales = log_scales.repeat(1, 1, self.n_sample, 1).reshape(means.shape[0], -1, dim, self.nr_mix) # [B, N*n_sample, dim, nr_mix]
39
40 means = self.activate_mean(means.type(torch.float32), mean_activation)
41 log_scales = log_scales.type(torch.float32)
42
43 logits = F.softmax(logits.type(torch.float32), dim=-1)
44
45 return logits, means, log_scales
46
47 def get_mix_params(self, logits, means, log_scales):
48 return means.squeeze(-1), log_scales.squeeze(-1), 1.0
49
50 def cdf_fn(self, x, means, log_scales):
51 """
52 Cumulative distribution function of the Gaussian distribution.
53 """
54 inv_std = torch.exp(-log_scales)
55 return 0.5 * (
56 1 + torch.erf((x - means) * inv_std / math.sqrt(2))
57 )
58
59 def log_pdf_fn(self, x, means, log_scales):
60 """
61 Log probability density function of the Gaussian distribution.
62 """
63 scales = torch.exp(log_scales)
64 var = scales**2
65 return (
66 -((x - means) ** 2) / (2 * var)

Callers 2

__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected