updated for readability
This commit is contained in:
@@ -1,12 +1,14 @@
|
||||
"""Definition of get_model function."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from io import BytesIO
|
||||
|
||||
import torch
|
||||
from minio import Minio
|
||||
|
||||
from .get import get
|
||||
from shared.utils import check_env
|
||||
|
||||
|
||||
def get_model(
|
||||
@@ -16,13 +18,24 @@ def get_model(
|
||||
"""Get model from model subfolder in bucket in Minio."""
|
||||
assert isinstance(client, Minio)
|
||||
assert isinstance(object_name, str)
|
||||
subfolder = 'models'
|
||||
object_name = f'{subfolder}/{object_name}'
|
||||
# get buffer
|
||||
buffer = get(
|
||||
client=client,
|
||||
object_name=object_name,
|
||||
)
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user