| 7 | |
| 8 | |
| 9 | class 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 |