MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / ComputeCovG

Class ComputeCovG

optimizers/kfac_utils.py:129–179  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

127
128
129class 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

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected