diff --git a/model/src/utils/__init__.py b/model/src/utils/__init__.py index ad58834..aa29844 100644 --- a/model/src/utils/__init__.py +++ b/model/src/utils/__init__.py @@ -1 +1 @@ -from .get_model_name import get_model_name +from .load_model import load_model diff --git a/model/src/utils/load_model.py b/model/src/utils/load_model.py new file mode 100644 index 0000000..332bea0 --- /dev/null +++ b/model/src/utils/load_model.py @@ -0,0 +1,40 @@ +"""Definition of load_model function.""" + +import logging +from pathlib import Path + +import torch +from minio import Minio +from models import VisualCommunicationModel + +from shared.data_store import get_model + +from .get_model_name import get_model_name + +DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + + +def load_model( + client: Minio, +) -> VisualCommunicationModel: + """Instantiate model with weights loaded from latest model saved in + MinIO.""" + assert isinstance(client, Minio) + # instantiate model + model = VisualCommunicationModel() + # get model object name + model_name_path = Path('model_name.txt') + model_object_name = get_model_name(path=model_name_path) + logging.info('using model: %s', model_object_name) + # load model from minio + model_checkpoint = get_model( + client=client, + object_name=model_object_name, + ) + model.load_state_dict(model_checkpoint) + # clean memory + model_checkpoint.clear() + # move model to selected device + model = model.to(DEVICE) + logging.debug('finished') + return model