From cbac5ded35e0db7dda404c0959639ebe01ea9f21 Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Sun, 21 Sep 2025 19:40:01 +0200 Subject: [PATCH] ruff format fixes --- .../repositories/annotation_repository.py | 20 +++++++--- scripts/query_ai_model.py | 38 +++++++++---------- 2 files changed, 34 insertions(+), 24 deletions(-) diff --git a/data_store/repositories/annotation_repository.py b/data_store/repositories/annotation_repository.py index b2974e8..2d6e582 100644 --- a/data_store/repositories/annotation_repository.py +++ b/data_store/repositories/annotation_repository.py @@ -60,17 +60,27 @@ class AnnotationRepository(MinioAdapter, AnnotationRepositoryInterface): annotation = Annotation( angle=AngleEnum(data["angle"]) if data["angle"] is not None else None, angle_uncertainty=bool(data["angle_uncertainty"]), - colour=ColourEnum(data["colour"]) if data["colour"] is not None else None, + colour=ColourEnum(data["colour"]) + if data["colour"] is not None + else None, colour_uncertainty=bool(data["colour_uncertainty"]), - contact=ContactEnum(data["contact"]) if data["contact"] is not None else None, + contact=ContactEnum(data["contact"]) + if data["contact"] is not None + else None, contact_uncertainty=bool(data["contact_uncertainty"]), depth=DepthEnum(data["depth"]) if data["depth"] is not None else None, depth_uncertainty=bool(data["depth_uncertainty"]), - distance=DistanceEnum(data["distance"]) if data["distance"] is not None else None, + distance=DistanceEnum(data["distance"]) + if data["distance"] is not None + else None, distance_uncertainty=bool(data["distance_uncertainty"]), - lighting=LightingEnum(data["lighting"]) if data["lighting"] is not None else None, + lighting=LightingEnum(data["lighting"]) + if data["lighting"] is not None + else None, lighting_uncertainty=bool(data["lighting_uncertainty"]), - point_of_view=PointOfViewEnum(data["point_of_view"]) if data["point_of_view"] is not None else None, + point_of_view=PointOfViewEnum(data["point_of_view"]) + if data["point_of_view"] is not None + else None, point_of_view_uncertainty=bool(data["point_of_view_uncertainty"]), ) except (KeyError, ValueError) as e: diff --git a/scripts/query_ai_model.py b/scripts/query_ai_model.py index dcac18e..46319a2 100644 --- a/scripts/query_ai_model.py +++ b/scripts/query_ai_model.py @@ -58,7 +58,7 @@ def query_model_with_image( prompt: str, allowed_answers: list[str], image: Image.Image, - model: str = "llava:13b" + model: str = "llava:13b", ) -> str: """ Query an Ollama instance with a prompt, image, and constrained answers. @@ -77,9 +77,9 @@ def query_model_with_image( # Read environment variables ollama_endpoint = str(os.getenv("OLLAMA_ENDPOINT")) # Handle RGBA images by converting to RGB - if image.mode == 'RGBA': + if image.mode == "RGBA": # Convert RGBA to RGB by compositing against white background - rgb_image = Image.new('RGB', image.size, (255, 255, 255)) + rgb_image = Image.new("RGB", image.size, (255, 255, 255)) rgb_image.paste(image, mask=image.split()[-1]) # Use alpha as mask image = rgb_image # Convert PIL Image to base64 @@ -97,14 +97,12 @@ def query_model_with_image( "model": model, "prompt": prompt_with_options, "images": [image_b64], - "stream": False + "stream": False, } # Make request to Ollama # N.B. the generate endpoint does not retain context between calls response = requests.post( - f"{ollama_endpoint}/api/generate", - json=payload, - timeout=60 + f"{ollama_endpoint}/api/generate", json=payload, timeout=60 ) response.raise_for_status() result = response.json() @@ -197,11 +195,15 @@ def calculate_model_prompt_accuracy( expected_answer = getattr(job.annotation, category.lower()) # Score the response score = score_response(response, expected_answer) - print(f"Job {job_id}: Response '{response}' - Expected answer '{expected_answer}' - Score: {score}") + print( + f"Job {job_id}: Response '{response}' - Expected answer '{expected_answer}' - Score: {score}" + ) total_score += score # Calculate accuracy jobs_evaluated = total_jobs - jobs_skipped - print(f"Total Jobs: {total_jobs}, Jobs Skipped: {jobs_skipped}, Jobs Evaluated: {jobs_evaluated}") + print( + f"Total Jobs: {total_jobs}, Jobs Skipped: {jobs_skipped}, Jobs Evaluated: {jobs_evaluated}" + ) if jobs_evaluated == 0: raise ValueError("No jobs were evaluated. All jobs were skipped.") accuracy = total_score / jobs_evaluated @@ -237,17 +239,11 @@ def refine_prompt( "Respond with only the refined prompt." ) # Prepare request payload - payload = { - "model": model, - "prompt": refinement_instruction, - "stream": False - } + payload = {"model": model, "prompt": refinement_instruction, "stream": False} # Make request to Ollama # N.B. the generate endpoint does not retain context between calls response = requests.post( - f"{ollama_endpoint}/api/generate", - json=payload, - timeout=60 + f"{ollama_endpoint}/api/generate", json=payload, timeout=60 ) response.raise_for_status() result = response.json() @@ -290,9 +286,13 @@ if __name__ == "__main__": print(f"Prompt accuracy for category '{category}': {accuracy:.2%}") # Stop if accuracy is 100% if accuracy == 1.0: - print(f"Achieved 100% accuracy for category '{category}'. Stopping refinement.") + print( + f"Achieved 100% accuracy for category '{category}'. Stopping refinement." + ) break - print(f"Final prompt accuracy map for category '{category}': {prompt_accuracy_map}") + print( + f"Final prompt accuracy map for category '{category}': {prompt_accuracy_map}" + ) category_events[category] = prompt_accuracy_map # save results to json file with open("prompt_accuracy_results.json", "w", encoding="utf-8") as f: