49 lines
1.3 KiB
Python
49 lines
1.3 KiB
Python
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,
|
|
)
|