From ba0fc050fed16c419f862318b36e677e146ac197 Mon Sep 17 00:00:00 2001 From: brian Date: Sat, 19 Oct 2024 21:37:08 +0000 Subject: [PATCH] fixed mypy types --- model/src/utils/get_class.py | 5 +++-- model/src/utils/load_model.py | 2 +- model/train.py | 4 ++-- other/transfer_images_mongo_minio.py | 4 ++-- shared/database/__init__.py | 3 ++- shared/database/utils/save_dataset.py | 6 ++++-- tests/generate_random_prediction_test.py | 4 ++-- tests/get_visual_communication_test.py | 4 ++-- tests/image_upload_test.py | 4 ++-- tests/image_upload_to_server_test.py | 5 ++--- tests/prediction_upload_test.py | 5 ++--- tests/total_annotated_test.py | 5 ++--- tests/total_documents_test.py | 5 ++--- web_ui/src/main.py | 4 ++-- 14 files changed, 30 insertions(+), 30 deletions(-) diff --git a/model/src/utils/get_class.py b/model/src/utils/get_class.py index 3b245b0..a4ac7da 100644 --- a/model/src/utils/get_class.py +++ b/model/src/utils/get_class.py @@ -1,8 +1,9 @@ from importlib import import_module -from types import ModuleType + +from torch.nn import Module -def get_class(path: str) -> ModuleType: +def get_class(path: str) -> Module: parts = path.split('.') module_path = '.'.join(parts[:-1]) class_name = parts[-1] diff --git a/model/src/utils/load_model.py b/model/src/utils/load_model.py index 912db6e..579421f 100644 --- a/model/src/utils/load_model.py +++ b/model/src/utils/load_model.py @@ -5,8 +5,8 @@ from pathlib import Path import torch from minio import Minio -from src.models import VisualCommunicationModel +from model.src.models import VisualCommunicationModel from shared.data_store import get_model from .get_model_name import get_model_name diff --git a/model/train.py b/model/train.py index 3855c19..8a13519 100644 --- a/model/train.py +++ b/model/train.py @@ -13,11 +13,11 @@ from ignite.handlers.param_scheduler import create_lr_scheduler_with_warmup from ignite.handlers.tensorboard_logger import TensorboardLogger from ignite.handlers.tqdm_logger import ProgressBar from ignite.metrics import Average, Loss, RunningAverage -from src.utils import VCDADataset, get_class from torch import nn from torch.optim.lr_scheduler import ExponentialLR from torch.utils.data import DataLoader +from model.src.utils import VCDADataset, get_class from shared.data_store import connect_minio @@ -90,7 +90,7 @@ val_metrics = { } evaluator = create_supervised_evaluator( model=model, - metrics=val_metrics, + metrics=val_metrics, # type: ignore device=device, ) ProgressBar(desc='Val', ncols=80).attach(evaluator) diff --git a/other/transfer_images_mongo_minio.py b/other/transfer_images_mongo_minio.py index 87ad971..0b9fcd3 100644 --- a/other/transfer_images_mongo_minio.py +++ b/other/transfer_images_mongo_minio.py @@ -11,7 +11,7 @@ from dotenv import load_dotenv from pymongo.collection import Collection from shared.data_store import connect_minio, put -from shared.database import VisualCommunication, connect +from shared.database import VisualCommunication, connect_mongodb from shared.utils import check_env, setup_logging @@ -82,7 +82,7 @@ if __name__ == '__main__': # connect to minIO minio_client = connect_minio() # connect to MongoDB - collection, db, client = connect() + collection, db, client = connect_mongodb() # list documents in mongoDB id_list = list_mongo_document_ids(collection) for doc_id in id_list: diff --git a/shared/database/__init__.py b/shared/database/__init__.py index a1321c1..ed187bd 100644 --- a/shared/database/__init__.py +++ b/shared/database/__init__.py @@ -1,10 +1,11 @@ """Database module content.""" + from __future__ import annotations from .classes.dataset import Dataset from .classes.exceptions import NoDocumentFoundException from .classes.visual_communication import VisualCommunication -from .utils.connect import connect +from .utils.connect_mongodb import connect_mongodb from .utils.count_documents import count_documents from .utils.get_visual_communication import get_visual_communication from .utils.list_names import list_names diff --git a/shared/database/utils/save_dataset.py b/shared/database/utils/save_dataset.py index f4605ec..0759152 100755 --- a/shared/database/utils/save_dataset.py +++ b/shared/database/utils/save_dataset.py @@ -20,10 +20,12 @@ def save_dataset( if __name__ == '__main__': from dotenv import load_dotenv + load_dotenv('local.env') - from shared.database import list_names, connect + from shared.database import connect_mongodb, list_names + # connect to database - collection, db, client = connect() + collection, db, client = connect_mongodb() print(client.server_info()) name_list = list_names(collection=collection, only_with_annotation=True) diff --git a/tests/generate_random_prediction_test.py b/tests/generate_random_prediction_test.py index 899ffe0..03dda23 100644 --- a/tests/generate_random_prediction_test.py +++ b/tests/generate_random_prediction_test.py @@ -5,7 +5,7 @@ from pathlib import Path from dotenv import load_dotenv from shared.data_store import connect_minio -from shared.database import connect as connect_mongo +from shared.database import connect_mongodb from shared.database.classes import VisualCommunication from shared.utils import check_env, setup_logging @@ -21,7 +21,7 @@ if __name__ == '__main__': # connect to minIO minio_client = connect_minio() # connect to MongoDB - collection, db, client = connect_mongo() + collection, db, client = connect_mongodb() # get list of image paths test_dir = Path(__file__).parent img_dir = test_dir / 'imgs' diff --git a/tests/get_visual_communication_test.py b/tests/get_visual_communication_test.py index 5964917..9d84316 100644 --- a/tests/get_visual_communication_test.py +++ b/tests/get_visual_communication_test.py @@ -5,7 +5,7 @@ from pathlib import Path from dotenv import load_dotenv from shared.data_store import connect_minio -from shared.database import connect, get_visual_communication +from shared.database import connect_mongodb, get_visual_communication from shared.utils import check_env, setup_logging if __name__ == '__main__': @@ -20,7 +20,7 @@ if __name__ == '__main__': # connect to minIO minio_client = connect_minio() # connect to MongoDB - collection, db, client = connect() + collection, db, client = connect_mongodb() # get visual communication vis_com = get_visual_communication(collection) print(repr(vis_com)) diff --git a/tests/image_upload_test.py b/tests/image_upload_test.py index c7a94bf..0b979e6 100644 --- a/tests/image_upload_test.py +++ b/tests/image_upload_test.py @@ -6,7 +6,7 @@ from dotenv import load_dotenv from pymongo.errors import DuplicateKeyError from shared.data_store import connect_minio -from shared.database import connect as connect_mongo +from shared.database import connect_mongodb from shared.database.classes import VisualCommunication from shared.utils import check_env, setup_logging @@ -22,7 +22,7 @@ if __name__ == '__main__': # connect to minIO minio_client = connect_minio() # connect to MongoDB - collection, db, client = connect_mongo() + collection, db, client = connect_mongodb() # get list of image paths test_dir = Path(__file__).parent img_dir = test_dir / 'imgs' diff --git a/tests/image_upload_to_server_test.py b/tests/image_upload_to_server_test.py index 6e2d59e..e3f353c 100644 --- a/tests/image_upload_to_server_test.py +++ b/tests/image_upload_to_server_test.py @@ -6,8 +6,7 @@ from dotenv import load_dotenv from pymongo.errors import DuplicateKeyError from shared.data_store import connect_minio -from shared.database import VisualCommunication -from shared.database import connect as connect_mongo +from shared.database import VisualCommunication, connect_mongodb from shared.utils import check_env, setup_logging if __name__ == '__main__': @@ -22,7 +21,7 @@ if __name__ == '__main__': # connect to minIO minio_client = connect_minio() # connect to MongoDB - collection, db, client = connect_mongo() + collection, db, client = connect_mongodb() # get list of image paths ext_img_dir = Path('/Volumes/BW-PSSD/Mixed Methods/') assert ext_img_dir.exists() diff --git a/tests/prediction_upload_test.py b/tests/prediction_upload_test.py index 9f635e7..45c9e5a 100644 --- a/tests/prediction_upload_test.py +++ b/tests/prediction_upload_test.py @@ -5,8 +5,7 @@ from pathlib import Path from dotenv import load_dotenv from shared.data_store import connect_minio -from shared.database import connect as connect_mongo -from shared.database import upsert_prediction +from shared.database import connect_mongodb, upsert_prediction from shared.database.classes import VisualCommunication from shared.utils import check_env, setup_logging @@ -22,7 +21,7 @@ if __name__ == '__main__': # connect to minIO minio_client = connect_minio() # connect to MongoDB - collection, db, client = connect_mongo() + collection, db, client = connect_mongodb() # get list of image paths test_dir = Path(__file__).parent img_dir = test_dir / 'imgs' diff --git a/tests/total_annotated_test.py b/tests/total_annotated_test.py index 8c94f27..12c41ad 100644 --- a/tests/total_annotated_test.py +++ b/tests/total_annotated_test.py @@ -5,8 +5,7 @@ from pathlib import Path from dotenv import load_dotenv -from shared.database import connect -from shared.database import count_documents +from shared.database import connect_mongodb, count_documents if __name__ == '__main__': # prepare env vars @@ -15,7 +14,7 @@ if __name__ == '__main__': load_dotenv(env_path) os.environ['MONGO_HOST'] = 'localhost' # connect to database - collection, db, client = connect() + collection, db, client = connect_mongodb() # get visual communication num_docs = count_documents( collection=collection, diff --git a/tests/total_documents_test.py b/tests/total_documents_test.py index 0242940..fffe1f4 100644 --- a/tests/total_documents_test.py +++ b/tests/total_documents_test.py @@ -5,8 +5,7 @@ from pathlib import Path from dotenv import load_dotenv -from shared.database import connect -from shared.database import count_documents +from shared.database import connect_mongodb, count_documents if __name__ == '__main__': # prepare env vars @@ -15,7 +14,7 @@ if __name__ == '__main__': load_dotenv(env_path) os.environ['MONGO_HOST'] = 'localhost' # connect to database - collection, db, client = connect() + collection, db, client = connect_mongodb() # get visual communication num_docs = count_documents( collection=collection, diff --git a/web_ui/src/main.py b/web_ui/src/main.py index 17c4f82..3d61313 100644 --- a/web_ui/src/main.py +++ b/web_ui/src/main.py @@ -5,7 +5,7 @@ from __future__ import annotations import os from shared.data_store import connect_minio -from shared.database import connect +from shared.database import connect_mongodb from shared.utils import check_env, setup_logging from .app import init_app @@ -17,7 +17,7 @@ check_env() setup_logging() # connect to database -collection, db, client = connect() +collection, db, client = connect_mongodb() # connect to minio minio_client = connect_minio()