Base class for all models
| 4 | |
| 5 | |
| 6 | class BaseModel(nn.Module): |
| 7 | """ |
| 8 | Base class for all models |
| 9 | """ |
| 10 | def __init__(self): |
| 11 | super(BaseModel, self).__init__() |
| 12 | self.logger = logging.getLogger(self.__class__.__name__) |
| 13 | |
| 14 | def forward(self, *input): |
| 15 | """ |
| 16 | Forward pass logic |
| 17 | |
| 18 | :return: Model output |
| 19 | """ |
| 20 | raise NotImplementedError |
| 21 | |
| 22 | def summary(self): |
| 23 | """ |
| 24 | Model summary |
| 25 | """ |
| 26 | model_parameters = filter(lambda p: p.requires_grad, self.parameters()) |
| 27 | params = sum([np.prod(p.size()) for p in model_parameters]) |
| 28 | self.logger.info('Trainable parameters: {}'.format(params)) |
| 29 | self.logger.info(self) |
nothing calls this directly
no outgoing calls
no test coverage detected