(self, model)
| 27 | self.register_buffer('num_updates', torch.tensor(0, dtype=torch.int)) |
| 28 | |
| 29 | def forward(self, model): |
| 30 | decay = self.decay |
| 31 | |
| 32 | if self.num_updates >= 0: |
| 33 | self.num_updates += 1 |
| 34 | decay = min(self.decay, (1 + self.num_updates) / (10 + self.num_updates)) |
| 35 | |
| 36 | one_minus_decay = 1.0 - decay |
| 37 | |
| 38 | with torch.no_grad(): |
| 39 | m_param = dict(model.named_parameters()) |
| 40 | shadow_params = dict(self.named_buffers()) |
| 41 | |
| 42 | for key in m_param: |
| 43 | if m_param[key].requires_grad: |
| 44 | sname = self.m_name2s_name[key] |
| 45 | shadow_params[sname] = shadow_params[sname].type_as(m_param[key]) |
| 46 | shadow_params[sname].sub_(one_minus_decay * (shadow_params[sname] - m_param[key])) |
| 47 | else: |
| 48 | assert not key in self.m_name2s_name |
| 49 | |
| 50 | def copy_to(self, model): |
| 51 | m_param = dict(model.named_parameters()) |
nothing calls this directly
no outgoing calls
no test coverage detected