diff --git a/tests/new_model_to_minio_test.py b/tests/new_model_to_minio_test.py new file mode 100644 index 0000000..33234e3 --- /dev/null +++ b/tests/new_model_to_minio_test.py @@ -0,0 +1,32 @@ +"""Definition of function to generate a new randomly initialized model and save +it in Minio datastore.""" + +from pathlib import Path + +from dotenv import load_dotenv +from torchinfo import summary + +from model.src.models import VisualCommunicationModel +from shared.data_store import connect_minio, put_model +from shared.utils import setup_logging + +if __name__ == '__main__': + # load in env file + env_path = Path(__file__).parent.parent / 'server.env' + assert env_path.exists() + load_dotenv(env_path) + # setup logging + setup_logging() + # connect to minio + client = connect_minio() + # instantiate model + model = VisualCommunicationModel(download_resnet_weights=True) + # show model weights + summary(model) + # put buffer in minio bucket + hash_str = put_model( + client=client, + model=model, + ) + print(f"hash string: {hash_str}") + print('saved to Minio')