44 lines
1.1 KiB
Python
44 lines
1.1 KiB
Python
"""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
|