(self, a)
| 2403 | return torch.einsum("...jk,...kl->...jl", Q, torch.transpose(V, -1, -2)) |
| 2404 | |
| 2405 | def eigh(self, a): |
| 2406 | return torch.linalg.eigh(a) |
| 2407 | |
| 2408 | def kl_div(self, p, q, mass=False, eps=1e-16): |
| 2409 | value = torch.sum(p * torch.log(p / q + eps)) |