From 31156eea54f36fe51856bf703a7f1dc8542c209a Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Mon, 22 Jul 2024 23:45:54 +0200 Subject: [PATCH] updated classes --- shared/data_store/get_model.py | 23 ++++++++++++++++++++--- 1 file changed, 20 insertions(+), 3 deletions(-) diff --git a/shared/data_store/get_model.py b/shared/data_store/get_model.py index 3610430..45f5bf6 100644 --- a/shared/data_store/get_model.py +++ b/shared/data_store/get_model.py @@ -1,14 +1,31 @@ """Definition of get_model function.""" +import logging +from collections import OrderedDict + +import torch from minio import Minio -from torch.nn import Module + +from .get import get def get_model( client: Minio, object_name: str, -) -> Module: +) -> OrderedDict: """Get model from model subfolder in bucket in Minio.""" assert isinstance(client, Minio) assert isinstance(object_name, str) - raise NotImplementedError() + subfolder = 'models' + object_name = f'{subfolder}/{object_name}' + # get buffer + buffer = get( + client=client, + object_name=object_name, + ) + logging.debug('buffer size: %s', buffer.getbuffer().nbytes) + # convert data to model checkpoint + buffer.seek(0) + model_content = torch.load(buffer) + logging.debug('finished') + return model_content