MCPcopy Create free account
hub / github.com/Amanda-Zheng/SFGC / ReparamModule

Class ReparamModule

models/reparam_module.py:9–150  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class ReparamModule(nn.Module):
10 def _get_module_from_name(self, mn):
11 if mn == '':
12 return self
13 m = self
14 for p in mn.split('.'):
15 m = getattr(m, p)
16 return m
17
18 def __init__(self, module):
19 super(ReparamModule, self).__init__()
20 self.module = module
21
22 param_infos = [] # (module name/path, param name)
23 shared_param_memo = {}
24 shared_param_infos = [] # (module name/path, param name, src module name/path, src param_name)
25 params = []
26 param_numels = []
27 param_shapes = []
28 for mn, m in self.named_modules():
29 for n, p in m.named_parameters(recurse=False):
30 if p is not None:
31 if p in shared_param_memo:
32 shared_mn, shared_n = shared_param_memo[p]
33 shared_param_infos.append((mn, n, shared_mn, shared_n))
34 else:
35 shared_param_memo[p] = (mn, n)
36 param_infos.append((mn, n))
37 params.append(p.detach())
38 param_numels.append(p.numel())
39 param_shapes.append(p.size())
40
41 assert len(set(p.dtype for p in params)) <= 1, \
42 "expects all parameters in module to have same dtype"
43
44 # store the info for unflatten
45 self._param_infos = tuple(param_infos)
46 self._shared_param_infos = tuple(shared_param_infos)
47 self._param_numels = tuple(param_numels)
48 self._param_shapes = tuple(param_shapes)
49
50 # flatten
51 flat_param = nn.Parameter(torch.cat([p.reshape(-1) for p in params], 0))
52 self.register_parameter('flat_param', flat_param)
53 self.param_numel = flat_param.numel()
54 del params
55 del shared_param_memo
56
57 # deregister the names as parameters
58 for mn, n in self._param_infos:
59 delattr(self._get_module_from_name(mn), n)
60 for mn, n, _, _ in self._shared_param_infos:
61 delattr(self._get_module_from_name(mn), n)
62
63 # register the views as plain attributes
64 self._unflatten_param(self.flat_param)
65
66 # now buffers

Callers 3

distillMethod · 0.90
distillMethod · 0.90
distillMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected