MCPcopy Create free account
hub / github.com/OpenGVLab/InternVL / __init__

Method __init__

classification/ema_deepspeed.py:13–30  ·  view source on GitHub ↗
(self, model, decay=0.9999, use_num_updates=True)

Source from the content-addressed store, hash-verified

11 """
12
13 def __init__(self, model, decay=0.9999, use_num_updates=True):
14 super().__init__()
15 if decay < 0.0 or decay > 1.0:
16 raise ValueError('Decay must be between 0 and 1')
17
18 self.m_name2s_name = {}
19 self.decay = decay
20 self.num_updates = 0 if use_num_updates else -1
21
22 with GatheredParameters(model.parameters(), fwd_module=self):
23 for name, p in model.named_parameters():
24 if p.requires_grad:
25 # remove as '.'-character is not allowed in buffers
26 s_name = name.replace('.', '')
27 self.m_name2s_name.update({name: s_name})
28 self.register_buffer(s_name, p.clone().detach().data)
29 # remove as '.'-character is not allowed in buffers
30 self.collected_params = []
31
32 def forward(self, model):
33 decay = self.decay

Callers

nothing calls this directly

Calls 1

updateMethod · 0.80

Tested by

no test coverage detected