| 127 | |
| 128 | |
| 129 | class ComputeCovG: |
| 130 | |
| 131 | @classmethod |
| 132 | def compute_cov_g(cls, g, layer, batch_averaged=False): |
| 133 | """ |
| 134 | :param g: gradient |
| 135 | :param layer: the corresponding layer |
| 136 | :param batch_averaged: if the gradient is already averaged with the batch size? |
| 137 | :return: |
| 138 | """ |
| 139 | # batch_size = g.size(0) |
| 140 | return cls.__call__(g, layer, batch_averaged) |
| 141 | |
| 142 | @classmethod |
| 143 | def __call__(cls, g, layer, batch_averaged): |
| 144 | if isinstance(layer, nn.Conv2d): |
| 145 | cov_g = cls.conv2d(g, layer, batch_averaged) |
| 146 | elif isinstance(layer, nn.Linear): |
| 147 | cov_g = cls.linear(g, layer, batch_averaged) |
| 148 | else: |
| 149 | cov_g = None |
| 150 | |
| 151 | return cov_g |
| 152 | |
| 153 | @staticmethod |
| 154 | def conv2d(g, layer, batch_averaged): |
| 155 | # g: batch_size * n_filters * out_h * out_w |
| 156 | # n_filters is actually the output dimension (analogous to Linear layer) |
| 157 | spatial_size = g.size(2) * g.size(3) |
| 158 | batch_size = g.shape[0] |
| 159 | g = g.transpose(1, 2).transpose(2, 3) |
| 160 | g = try_contiguous(g) |
| 161 | g = g.view(-1, g.size(-1)) |
| 162 | |
| 163 | if batch_averaged: |
| 164 | g = g * batch_size |
| 165 | g = g * spatial_size |
| 166 | cov_g = g.t() @ (g / g.size(0)) |
| 167 | |
| 168 | return cov_g |
| 169 | |
| 170 | @staticmethod |
| 171 | def linear(g, layer, batch_averaged): |
| 172 | # g: batch_size * out_dim |
| 173 | batch_size = g.size(0) |
| 174 | |
| 175 | if batch_averaged: |
| 176 | cov_g = g.t() @ (g * batch_size) |
| 177 | else: |
| 178 | cov_g = g.t() @ (g / batch_size) |
| 179 | return cov_g |