"""Definition of get_model function.""" import logging import os from collections import OrderedDict from io import BytesIO import torch from minio import Minio from shared.utils import check_env def get_model( client: Minio, object_name: str, ) -> OrderedDict: """Get model from model subfolder in bucket in Minio.""" assert isinstance(client, Minio) assert isinstance(object_name, str) assert len(object_name) > 0 # prepare arguments check_env({'MINIO_BUCKET_NAME_MODELS'}) bucket_name = str(os.getenv('MINIO_BUCKET_NAME_MODELS')) object_name = f'models/{object_name}' # get object from bucket try: response = client.get_object( bucket_name=bucket_name, object_name=object_name, ) buffer = BytesIO(response.data) except Exception as exc: logging.error('failed getting data from MinIO') raise exc finally: response.close() response.release_conn() # convert data to model checkpoint buffer.seek(0) model_content = torch.load(buffer) logging.debug('finished') return model_content