-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathBaseModel.py
More file actions
36 lines (29 loc) · 1.14 KB
/
Copy pathBaseModel.py
File metadata and controls
36 lines (29 loc) · 1.14 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
import os
import utils
import torch
from torch import nn
from pathlib import Path
from abc import ABC, abstractmethod
class BaseModel(nn.Module, ABC):
def __init__(self):
super(BaseModel, self).__init__()
self.res_dir = os.path.join(utils.get_res_path(), self.TRAINED_MODELS_DIR)
self.device = utils.get_device()
@property
@abstractmethod
def MODEL_NAME(self):
...
@property
@abstractmethod
def TRAINED_MODELS_DIR(self):
...
def save_model(self, res_dir=None, model_name=None):
model_name = self.MODEL_NAME if model_name is None else model_name
res_dir = self.res_dir if res_dir is None else res_dir
Path(res_dir).mkdir(exist_ok=True, parents=True)
torch.save(self.state_dict(), os.path.join(res_dir, model_name))
def load_model(self, res_dir=None, model_name=None):
model_name = self.MODEL_NAME if model_name is None else model_name
res_dir = self.res_dir if res_dir is None else res_dir
print(f'Loading model {model_name}...')
self.load_state_dict(torch.load(os.path.join(res_dir, model_name), map_location=self.device))