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

Class EMA

prototype/utils/ema.py:6–83  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4
5
6class EMA(object):
7 def __init__(self, model, decay, copy_init=False, use_double=False, inner_T=1, warmup=1):
8 self.ema_state_dict = OrderedDict()
9 self.logger = get_logger(__name__)
10 self.logger.info(f'EMA: decay={decay}, copy_init={copy_init}, \
11 use_double={use_double}, inner_T={inner_T}, warmup={warmup}')
12 self.use_double = use_double
13 self.inner_T = inner_T
14 self.decay = decay
15 self.warmup = warmup
16 if self.inner_T > 1:
17 self.decay = self.decay ** self.inner_T
18 self.logger.info('EMA: effective decay={}'.format(self.decay))
19 state_dict = model.state_dict()
20 if copy_init:
21 if self.use_double:
22 for k, v in state_dict.items():
23 self.ema_state_dict[k] = v.data.clone().double()
24 else:
25 for k, v in state_dict.items():
26 self.ema_state_dict[k] = v.data.clone().float()
27 else:
28 if self.use_double:
29 for k, v in state_dict.items():
30 self.ema_state_dict[k] = torch.zeros_like(v).double()
31 else:
32 for k, v in state_dict.items():
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():
54 tmp = v.data.clone()
55 v.data.copy_(self.ema_state_dict[k].data)
56 self.ema_state_dict[k].data.copy_(tmp)
57
58 def recover(self, model):
59 state_dict = model.state_dict()
60 for k, v in self.ema_state_dict.items():
61 tmp = v.data.clone()
62 v.data.copy_(state_dict[k].data)
63 state_dict[k].data.copy_(tmp)

Callers 5

build_optimizerMethod · 0.90
build_optimizerMethod · 0.90
build_optimizerMethod · 0.90
build_optimizerMethod · 0.90
build_optimizerMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected