From 0fbdc3865c234d8b6ef4ba596503351442a8ceb6 Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Tue, 21 May 2024 19:15:44 +0200 Subject: [PATCH] updated main script --- model/src/main.py | 45 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 45 insertions(+) create mode 100644 model/src/main.py diff --git a/model/src/main.py b/model/src/main.py new file mode 100644 index 0000000..b80361c --- /dev/null +++ b/model/src/main.py @@ -0,0 +1,45 @@ +"""Main script to be run by service.""" +from __future__ import annotations + +from hashlib import md5 +from io import BytesIO + +import torch +from min_io import connect +from min_io import put +from models import VisualCommunicationModel +from torchinfo import summary +from utils import check_env + +from shared.utils import setup_logging + + +DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + +if __name__ == '__main__': + # check that environment variables are set + check_env() + # setup logging + setup_logging() + # instantiate model + model = VisualCommunicationModel().to(DEVICE) + # show model weights + summary(model) + # for param in model.parameters(): + # logging.info(param.data) + # save model to buffer + buffer = BytesIO() + torch.save(model.state_dict(), buffer) + print(f"buffer size: {len(buffer.getvalue())}") + # calculate buffer hash + hash_str = md5(buffer.getbuffer()).hexdigest() + print(f"hash string: {hash_str}") + # connect to minio + client = connect() + # put buffer in minio bucket + put( + client=client, + buffer=buffer, + object_name=hash_str, + ) + print('saved to Minio')