| 7 | eps = 1e-3 |
| 8 | |
| 9 | class 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) |