From 01ee4664ffe23730958bb775bfda1f6ab5759b7a Mon Sep 17 00:00:00 2001 From: brian Date: Sat, 16 Nov 2024 18:18:13 +0000 Subject: [PATCH] implemented base functions --- shared/datastore/src/get_image.py | 21 ++++++------------ shared/datastore/src/get_model.py | 25 ++++++++------------- shared/datastore/src/put_image.py | 7 +++--- shared/datastore/src/put_model.py | 36 +++++++++++++++---------------- 4 files changed, 36 insertions(+), 53 deletions(-) diff --git a/shared/datastore/src/get_image.py b/shared/datastore/src/get_image.py index ef81762..596226b 100644 --- a/shared/datastore/src/get_image.py +++ b/shared/datastore/src/get_image.py @@ -2,13 +2,14 @@ import logging import os -from io import BytesIO from minio import Minio from PIL import Image from shared.utils import check_env +from .get import get + def get_image( client: Minio, @@ -23,19 +24,11 @@ def get_image( bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) object_name = f'images/{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() - buffer.seek(0) + buffer = get( + client=client, + bucket_name=bucket_name, + object_name=object_name, + ) # convert data to image image = Image.open(buffer) logging.debug('got data from %s', object_name) diff --git a/shared/datastore/src/get_model.py b/shared/datastore/src/get_model.py index 2cc40e7..59c349a 100644 --- a/shared/datastore/src/get_model.py +++ b/shared/datastore/src/get_model.py @@ -3,13 +3,14 @@ import logging import os from collections import OrderedDict -from io import BytesIO import torch from minio import Minio from shared.utils import check_env +from .get import get + def get_model( client: Minio, @@ -20,24 +21,16 @@ def get_model( 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')) + check_env({'MINIO_BUCKET_NAME'}) + bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) 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() + buffer = get( + client=client, + bucket_name=bucket_name, + object_name=object_name, + ) # convert data to model checkpoint - buffer.seek(0) model_content = torch.load(buffer) logging.debug('finished') return model_content diff --git a/shared/datastore/src/put_image.py b/shared/datastore/src/put_image.py index b5de63c..bea9dd6 100644 --- a/shared/datastore/src/put_image.py +++ b/shared/datastore/src/put_image.py @@ -22,14 +22,13 @@ def put_image( # get bucket name from env bucket_name = os.getenv('MINIO_BUCKET_NAME', default='') assert len(bucket_name) > 0 - # save image to buffer + # save data to buffer buffer = BytesIO() image.save(buffer, 'png') # get md5 of buffer checksum = md5(buffer.getbuffer()).hexdigest() - # determine object name - subfolder = 'images' - object_name = f'{subfolder}/{checksum}' + # set object name + object_name = f'images/{checksum}' # send data to bucket put( client=client, diff --git a/shared/datastore/src/put_model.py b/shared/datastore/src/put_model.py index 164b499..69436a9 100644 --- a/shared/datastore/src/put_model.py +++ b/shared/datastore/src/put_model.py @@ -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