MCPcopy Create free account
hub / github.com/MaureenZOU/TSAM / BaseModel

Class BaseModel

src/base/base_model.py:6–29  ·  view source on GitHub ↗

Base class for all models

Source from the content-addressed store, hash-verified

4
5
6class 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)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected