Files
visual_critical_discourse_a…/shared/datastore/src/get_model.py
T
2024-10-25 17:18:53 +00:00

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