(a, layer)
| 119 | |
| 120 | @staticmethod |
| 121 | def linear(a, layer): |
| 122 | # a: batch_size * in_dim |
| 123 | batch_size = a.size(0) |
| 124 | if layer.bias is not None: |
| 125 | a = torch.cat([a, a.new(a.size(0), 1).fill_(1)], 1) |
| 126 | return a.t() @ (a / batch_size) |
| 127 | |
| 128 | |
| 129 | class ComputeCovG: |