From e3e836c9fa8fc545e3258abe7942caf9827b6f73 Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Tue, 7 May 2024 21:22:25 +0200 Subject: [PATCH] fixed imports --- image_download/get_images.py | 3 +-- tests/model_outputs_from_annotation_test.py | 5 ++++- tests/prediction_upload_test.py | 4 ++-- tests/total_annotated_test.py | 5 +++-- tests/total_documents_test.py | 5 +++-- 5 files changed, 13 insertions(+), 9 deletions(-) diff --git a/image_download/get_images.py b/image_download/get_images.py index 90b6861..f167edb 100644 --- a/image_download/get_images.py +++ b/image_download/get_images.py @@ -6,10 +6,9 @@ from pathlib import Path import pandas as pd import requests +from classes import Instagram from retry import retry -from image_download.classes import Instagram - def get_sources() -> pd.DataFrame: """Get sources dateframe.""" diff --git a/tests/model_outputs_from_annotation_test.py b/tests/model_outputs_from_annotation_test.py index c4844fb..5e002fc 100644 --- a/tests/model_outputs_from_annotation_test.py +++ b/tests/model_outputs_from_annotation_test.py @@ -11,6 +11,9 @@ if __name__ == '__main__': in range(3) ] # generate random predictions - [vis_com.generate_random_prediction() for vis_com in vis_com_list] + [ + vis_com.generate_random_prediction() + for vis_com in vis_com_list + ] # type: ignore for vis_com in vis_com_list: print(vis_com) diff --git a/tests/prediction_upload_test.py b/tests/prediction_upload_test.py index 94246c0..7d8987b 100644 --- a/tests/prediction_upload_test.py +++ b/tests/prediction_upload_test.py @@ -7,7 +7,7 @@ from pathlib import Path from dotenv import load_dotenv from core.database import connect -from core.database import upsert_predictions +from core.database import upsert_prediction from core.database.classes import VisualCommunication if __name__ == '__main__': @@ -45,7 +45,7 @@ if __name__ == '__main__': for vis_com in vis_com_list: if vis_com.prediction is None: continue - upsert_predictions( + upsert_prediction( collection=collection, vis_com_name=vis_com.name, predictions=vis_com.prediction, diff --git a/tests/total_annotated_test.py b/tests/total_annotated_test.py index 0dafadb..0fe3717 100644 --- a/tests/total_annotated_test.py +++ b/tests/total_annotated_test.py @@ -6,7 +6,7 @@ from pathlib import Path from dotenv import load_dotenv from core.database import connect -from core.database import total_documents +from core.database import count_documents if __name__ == '__main__': # prepare env vars @@ -17,7 +17,8 @@ if __name__ == '__main__': # connect to database collection, db, client = connect() # get visual communication - num_docs = total_documents( + num_docs = count_documents( collection=collection, + only_with_annotation=True, ) print(f"number of annotated documents in database: {num_docs}") diff --git a/tests/total_documents_test.py b/tests/total_documents_test.py index e6e04a7..d7955fb 100644 --- a/tests/total_documents_test.py +++ b/tests/total_documents_test.py @@ -6,7 +6,7 @@ from pathlib import Path from dotenv import load_dotenv from core.database import connect -from core.database import total_documents +from core.database import count_documents if __name__ == '__main__': # prepare env vars @@ -17,7 +17,8 @@ if __name__ == '__main__': # connect to database collection, db, client = connect() # get visual communication - num_docs = total_documents( + num_docs = count_documents( collection=collection, + only_with_annotation=False, ) print(f"total number of documents in database: {num_docs}")