updated tests

This commit is contained in:
Brian Bjarke Jensen
2024-02-24 22:43:56 +01:00
parent d3eac22706
commit c3ad732448
3 changed files with 58 additions and 1 deletions
+1 -1
View File
@@ -18,4 +18,4 @@ if __name__ == "__main__":
print(client.server_info()) print(client.server_info())
# get visual communication # get visual communication
vis_com = get_visual_communication(collection) vis_com = get_visual_communication(collection)
print(vis_com.image) print(vis_com)
@@ -0,0 +1,13 @@
from src.database import ModelOutputs
if __name__ == "__main__":
# instantiate data object
annotation = {
}
vis_com_list = [ModelOutputs.from_annotation(path) for path in img_path_list]
# generate random predictions
[vis_com.generate_random_prediction() for vis_com in vis_com_list]
for vis_com in vis_com_list:
print(vis_com)
+44
View File
@@ -0,0 +1,44 @@
from pathlib import Path
from dotenv import load_dotenv
import os
import logging
from src.database import (
VisualCommunication,
connect,
upsert_predictions
)
if __name__ == "__main__":
# setup logging
fmt = (
'%(asctime)s | '
'%(levelname)s | '
'%(filename)s | '
'%(funcName)s | '
'%(message)s'
)
datefmt = '%Y-%m-%d %H:%M:%S'
logging.basicConfig(format=fmt, datefmt=datefmt, level=logging.INFO)
# get list of image paths
test_dir = Path(__file__).parent
img_dir = test_dir / "imgs"
img_path_list = [path for path in img_dir.glob("*.jpeg") if path.is_file()]
# instantiate data object
vis_com_list = [VisualCommunication.from_file(path) for path in img_path_list]
# generate random predictions
[vis_com.generate_random_prediction() for vis_com in vis_com_list]
# prepare env vars
env_path = test_dir.parent / "local.env"
assert env_path.exists()
load_dotenv(env_path)
os.environ["MONGO_HOST"] = "localhost"
# connect to database
collection, db, client = connect()
# upload visual communication
for vis_com in vis_com_list:
upsert_predictions(
collection=collection,
vis_com_name=vis_com.name,
predictions=vis_com.prediction,
)