(self, data, dataHat)
| 53 | return (locs, scales, weights) |
| 54 | |
| 55 | def loss(self, data, dataHat): |
| 56 | locs, scales, weights = dataHat |
| 57 | log_probs = -0.5 * T.sum( |
| 58 | (data.unsqueeze(-2) - locs).pow(2) / scales.pow(2) + |
| 59 | 2 * T.log(scales) + |
| 60 | T.log(T.tensor(2 * T.pi)), |
| 61 | dim=-1 |
| 62 | ) |
| 63 | log_weights = F.log_softmax(weights, dim=-1) |
| 64 | return -T.logsumexp(log_weights + log_probs, dim=-1) |
| 65 | |
| 66 | |
| 67 | def temp_sample(self, orig_pdist, temp): |
nothing calls this directly
no outgoing calls
no test coverage detected