MCPcopy Create free account
hub / github.com/Sense-GVT/DeCLIP / step

Method step

prototype/utils/ema.py:35–50  ·  view source on GitHub ↗
(self, model, curr_step=None)

Source from the content-addressed store, hash-verified

33 self.ema_state_dict[k] = torch.zeros_like(v).float()
34
35 def step(self, model, curr_step=None):
36 if curr_step is None:
37 decay = self.decay
38 else:
39 decay = min(self.decay, (1+curr_step)/(self.warmup+curr_step))
40
41 if curr_step % self.inner_T != 0:
42 return
43
44 state_dict = model.state_dict()
45 if self.use_double:
46 for k, v in state_dict.items():
47 self.ema_state_dict[k].mul_(decay).add_(1-decay, v.double())
48 else:
49 for k, v in state_dict.items():
50 self.ema_state_dict[k].mul_(decay).add_(1-decay, v.float())
51
52 def load_ema(self, model):
53 for k, v in model.state_dict().items():

Callers 5

trainMethod · 0.45
trainMethod · 0.45
trainMethod · 0.45
trainMethod · 0.45
trainMethod · 0.45

Calls 1

state_dictMethod · 0.80

Tested by

no test coverage detected