"""Definition of load_model function.""" import logging from pathlib import Path import torch from minio import Minio from model.src.models import VisualCommunicationModel from shared.datastore 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