| 10 | |
| 11 | |
| 12 | class BaseModel(torch.nn.Module): |
| 13 | def __init__(self): |
| 14 | super().__init__() |
| 15 | |
| 16 | def print_architecture(self, verbose=False): |
| 17 | name = type(self).__name__ |
| 18 | result = '-------------------%s---------------------\n' % name |
| 19 | total_num_params = 0 |
| 20 | for i, (name, child) in enumerate(self.named_children()): |
| 21 | if 'loss' in name: |
| 22 | continue |
| 23 | num_params = sum([p.numel() for p in child.parameters()]) |
| 24 | total_num_params += num_params |
| 25 | if verbose: |
| 26 | result += "%s: %3.3fM\n" % (name, (num_params / 1e6)) |
| 27 | for i, (name, grandchild) in enumerate(child.named_children()): |
| 28 | num_params = sum([p.numel() for p in grandchild.parameters()]) |
| 29 | if verbose: |
| 30 | result += "\t%s: %3.3fM\n" % (name, (num_params / 1e6)) |
| 31 | result += '[Network %s] Total number of parameters : %.3f M\n' % (name, total_num_params / 1e6) |
| 32 | result += '-----------------------------------------------\n' |
| 33 | print(result) |
| 34 | |
| 35 | def set_requires_grad(self, requires_grad): |
| 36 | for param in self.parameters(): |
| 37 | param.requires_grad = requires_grad |
| 38 | |
| 39 | def get_parameters_for_train(self): |
| 40 | return self.parameters() |
| 41 | |
| 42 | def forward(self): |
| 43 | raise NotImplementedError() |
| 44 | |
| 45 | |
| 46 |
nothing calls this directly
no outgoing calls
no test coverage detected