| 6 | |
| 7 | |
| 8 | class KFACOptimizer(optim.Optimizer): |
| 9 | def __init__(self, |
| 10 | model, |
| 11 | lr=0.001, |
| 12 | momentum=0.9, |
| 13 | stat_decay=0.95, |
| 14 | damping=0.001, |
| 15 | kl_clip=0.001, |
| 16 | weight_decay=0, |
| 17 | TCov=10, |
| 18 | TInv=100, |
| 19 | batch_averaged=True): |
| 20 | if lr < 0.0: |
| 21 | raise ValueError("Invalid learning rate: {}".format(lr)) |
| 22 | if momentum < 0.0: |
| 23 | raise ValueError("Invalid momentum value: {}".format(momentum)) |
| 24 | if weight_decay < 0.0: |
| 25 | raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) |
| 26 | defaults = dict(lr=lr, momentum=momentum, damping=damping, |
| 27 | weight_decay=weight_decay) |
| 28 | # TODO (CW): KFAC optimizer now only support model as input |
| 29 | super(KFACOptimizer, self).__init__(model.parameters(), defaults) |
| 30 | self.CovAHandler = ComputeCovA() |
| 31 | self.CovGHandler = ComputeCovG() |
| 32 | self.batch_averaged = batch_averaged |
| 33 | |
| 34 | self.known_modules = {'Linear', 'Conv2d'} |
| 35 | |
| 36 | self.modules = [] |
| 37 | self.grad_outputs = {} |
| 38 | |
| 39 | self.model = model |
| 40 | self._prepare_model() |
| 41 | |
| 42 | self.steps = 0 |
| 43 | |
| 44 | self.m_aa, self.m_gg = {}, {} |
| 45 | self.Q_a, self.Q_g = {}, {} |
| 46 | self.d_a, self.d_g = {}, {} |
| 47 | self.stat_decay = stat_decay |
| 48 | |
| 49 | self.kl_clip = kl_clip |
| 50 | self.TCov = TCov |
| 51 | self.TInv = TInv |
| 52 | |
| 53 | def _save_input(self, module, input): |
| 54 | if torch.is_grad_enabled() and self.steps % self.TCov == 0: |
| 55 | aa = self.CovAHandler(input[0].data, module) |
| 56 | # Initialize buffers |
| 57 | if self.steps == 0: |
| 58 | self.m_aa[module] = torch.diag(aa.new(aa.size(0)).fill_(1)) |
| 59 | update_running_stat(aa, self.m_aa[module], self.stat_decay) |
| 60 | |
| 61 | def _save_grad_output(self, module, grad_input, grad_output): |
| 62 | # Accumulate statistics for Fisher matrices |
| 63 | if self.steps % self.TCov == 0: |
| 64 | gg = self.CovGHandler(grad_output[0].data, module, self.batch_averaged) |
| 65 | # Initialize buffers |
nothing calls this directly
no outgoing calls
no test coverage detected