added function to load model
This commit is contained in:
@@ -1 +1 @@
|
|||||||
from .get_model_name import get_model_name
|
from .load_model import load_model
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user