implemented base functions
This commit is contained in:
@@ -9,6 +9,10 @@ import torch
|
||||
from minio import Minio
|
||||
from torch.nn import Module
|
||||
|
||||
from shared.utils import check_env
|
||||
|
||||
from .put import put
|
||||
|
||||
|
||||
def put_model(
|
||||
client: Minio,
|
||||
@@ -18,28 +22,22 @@ def put_model(
|
||||
used as object name."""
|
||||
assert isinstance(client, Minio)
|
||||
assert isinstance(model, Module)
|
||||
bucket_name = os.getenv('MINIO_BUCKET_NAME', default='')
|
||||
assert len(bucket_name) > 0
|
||||
subfolder = 'models'
|
||||
# save image to buffer
|
||||
# get bucket name from env
|
||||
check_env({'MINIO_BUCKET_NAME'})
|
||||
bucket_name = str(os.getenv('MINIO_BUCKET_NAME'))
|
||||
# save data to buffer
|
||||
buffer = BytesIO()
|
||||
torch.save(model.state_dict(), buffer)
|
||||
# get md5 of image
|
||||
checksum = md5(buffer.getbuffer()).hexdigest()
|
||||
# prepare for saving
|
||||
num_bytes = buffer.tell()
|
||||
buffer.seek(0)
|
||||
# set object name
|
||||
object_name = f'models/{checksum}'
|
||||
# send data to bucket
|
||||
object_name = f'{subfolder}/{checksum}'
|
||||
try:
|
||||
client.put_object(
|
||||
bucket_name=bucket_name,
|
||||
object_name=object_name,
|
||||
length=num_bytes,
|
||||
data=buffer,
|
||||
)
|
||||
except Exception as exc:
|
||||
logging.error('failed saving data to MinIO')
|
||||
raise exc
|
||||
logging.debug('saved data to %s', object_name)
|
||||
put(
|
||||
client=client,
|
||||
buffer=buffer,
|
||||
bucket_name=bucket_name,
|
||||
object_name=object_name,
|
||||
)
|
||||
logging.debug('finished')
|
||||
return checksum
|
||||
|
||||
Reference in New Issue
Block a user