p(x) = sum_i w[i] N(mu[i], sigma[i]^2 * I) config: w: shape K X 1, mixture coefficients, must sum to 1 mu: shape K X D, mean sigma: shape K X D, (diagonal) variance
(self, w, mu, sigma)
| 22 | """ |
| 23 | |
| 24 | def __init__(self, w, mu, sigma): |
| 25 | """ |
| 26 | p(x) = sum_i w[i] N(mu[i], sigma[i]^2 * I) |
| 27 | |
| 28 | config: |
| 29 | w: shape K X 1, mixture coefficients, must sum to 1 |
| 30 | mu: shape K X D, mean |
| 31 | sigma: shape K X D, (diagonal) variance |
| 32 | """ |
| 33 | super().__init__() |
| 34 | self.register_buffer('w', w) |
| 35 | self.register_buffer('mu', mu) |
| 36 | self.register_buffer('sigma', sigma) |
| 37 | self.K = w.shape[0] |
| 38 | self.D = mu.shape[1] |
| 39 | |
| 40 | @torch.no_grad() |
| 41 | def log_gaussian(self, x, mu, sigma): |
nothing calls this directly
no outgoing calls
no test coverage detected