MCPcopy Create free account
hub / github.com/VisionXLab/OF-Diff / forward

Method forward

ldm/modules/ema.py:29–48  ·  view source on GitHub ↗
(self, model)

Source from the content-addressed store, hash-verified

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())

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected