Files
visual-semiotic-ai-analysis/scripts/query_ai_model.py
T
Brian Bjarke Jensen 007161a927
Formatting Check / formatting-check (push) Successful in 9s
Python Code Quality / python-code-quality (push) Successful in 18s
Python Test / python-test (push) Successful in 30s
Sync Label Studio Annotations / sync_storage (push) Successful in 57s
code quality fixes
2025-09-26 00:02:02 +02:00

328 lines
11 KiB
Python

"""Script to query an AI model with a prompt and image and print the response."""
import os
import base64
import logging
import json
from io import BytesIO
import requests
import structlog
from dotenv import load_dotenv
from PIL import Image
from datetime import datetime
import mlflow
from python_utils import check_env
from data_store.repositories import JobRepository
from data_store.dto import (
AngleEnum,
ColourEnum,
ContactEnum,
DepthEnum,
DistanceEnum,
LightingEnum,
PointOfViewEnum,
)
RUN_NAME = f"query-ai-model_{datetime.now().strftime('%Y%m%d-%H%M%S')}"
MODEL = "llava:13b"
CATEGORY_PROMPT_MAP = {
"angle": (
"In the context of visual angle in the Kress and van Leeuwen framework, "
"analyze the image to determine the angle of the subject."
),
"colour": (
"In the context of colour in the Kress and van Leeuwen framework, "
"analyze the image to determine the colour saturation of the whole image."
),
"contact": (
"In the context of contact in the Kress and van Leeuwen framework, "
"analyze the image to determine the contact between subject and viewer."
),
"depth": (
"In the context of depth in the Kress and van Leeuwen framework, "
"analyze the image to determine the depth of the whole image."
),
"distance": (
"In the context of spatial distance in the Kress and van Leeuwen framework, "
"analyze the image to determine the spatial distance that the subject is portrayed at."
),
"lighting": (
"In the context of lighting in the Kress and van Leeuwen framework, "
"analyze the image to determine the lighting of the whole image."
),
"point_of_view": (
"In the context of point of view in the Kress and van Leeuwen framework, "
"analyze the image to determine the orientation of the subject in relation to the viewer."
),
}
def query_model_with_image(
prompt: str,
allowed_answers: list[str],
image: Image.Image,
model: str = "llava:13b",
) -> str:
"""
Query an Ollama instance with a prompt, image, and constrained answers.
Args:
prompt: The text prompt to send
allowed_answers: List of valid response options
image: PIL Image to analyze
model: The Ollama model to use (default: "llava:13b")
Returns:
The model's response text
"""
# Ensure required env vars are set
check_env("OLLAMA_ENDPOINT")
# Read environment variables
ollama_endpoint = str(os.getenv("OLLAMA_ENDPOINT"))
# Handle RGBA images by converting to RGB
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.paste(image, mask=image.split()[-1]) # Use alpha as mask
image = rgb_image
# Convert PIL Image to base64
buffer = BytesIO()
image.save(buffer, format="JPEG")
image_b64 = base64.b64encode(buffer.getvalue()).decode()
# Build prompt with allowed answers
prompt_with_options = (
f"{prompt}\n"
f"Please choose one of the following options: {', '.join(allowed_answers)}.\n"
"Respond with only the chosen option."
)
# Prepare request payload
payload = {
"model": model,
"prompt": prompt_with_options,
"images": [image_b64],
"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
)
response.raise_for_status()
result = response.json()
return str(result.get("response", "").strip())
def score_response(response: str, expected_answer: str) -> int:
"""
Score the model's response based on allowed answers.
Args:
response: The model's response text
expected_answer: The expected answer
Returns:
1 if the response is in allowed answers, else 0
"""
# Score response
return 1 if response.lower() == expected_answer.lower() else 0
def allowed_answers_for_category(
category: str,
) -> list[str]:
"""
Get allowed answers based on the category.
Args:
category: The category to determine allowed answers
Returns:
List of allowed answers
"""
match category.lower():
case "angle":
return [e.value for e in AngleEnum]
case "colour":
return [e.value for e in ColourEnum]
case "contact":
return [e.value for e in ContactEnum]
case "depth":
return [e.value for e in DepthEnum]
case "distance":
return [e.value for e in DistanceEnum]
case "lighting":
return [e.value for e in LightingEnum]
case "point_of_view":
return [e.value for e in PointOfViewEnum]
case _:
raise ValueError(f"Invalid category: {category}")
def calculate_model_prompt_accuracy(
model: str,
prompt: str,
category: str,
) -> float:
"""
Calculate the accuracy of the prompt against the expected answers for the given category.
"""
# Determine allowed answers based on category
allowed_answers = allowed_answers_for_category(category)
# Score each job in the repository)
jobs_skipped = 0
total_score = 0
with JobRepository() as job_repo:
job_ids = job_repo.list_all()
if not job_ids:
raise ValueError("No jobs found in JobRepository.")
total_jobs = len(job_ids)
# Get all responses for the prompt
for job_id in job_ids:
job = job_repo.get(job_id)
if not job:
jobs_skipped += 1
continue
# skip job if uncertainty flag is set for the category
if getattr(job.annotation, f"{category.lower()}_uncertainty"):
print(f"Job {job_id}: Skipping due to uncertainty flag.")
jobs_skipped += 1
continue
response = query_model_with_image(
model=model,
prompt=prompt,
allowed_answers=allowed_answers,
image=job.image,
)
# Get expected answer from job annotation
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}"
)
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}"
)
if jobs_evaluated == 0:
raise ValueError("No jobs were evaluated. All jobs were skipped.")
accuracy = total_score / jobs_evaluated
return accuracy
def refine_prompt(
model: str,
prompt_accuracy_map: dict[str, float],
) -> str:
"""
Refine the prompt to improve accuracy.
Args:
model: The Ollama model to use
prompt_accuracy_map: A mapping of prompts to their accuracies
Returns:
The refined prompt text
"""
# Ensure required env vars are set
check_env("OLLAMA_ENDPOINT")
# Read environment variables
ollama_endpoint = str(os.getenv("OLLAMA_ENDPOINT"))
# Convert prompt accuracy map to json string for better readability
prompt_accuracy_json = json.dumps(prompt_accuracy_map, indent=2)
# Build refinement instruction
refinement_instruction = (
"Given the following prompts and accuracies:\n"
f"{prompt_accuracy_json}\n"
"Please refine the prompt to improve the accuracy.\n"
"The refined prompt should be clear, concise, and focused on obtaining accurate responses.\n"
"Respond with only the refined prompt."
)
# Prepare request payload
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
)
response.raise_for_status()
result = response.json()
return str(result.get("response", "").strip())
if __name__ == "__main__":
load_dotenv()
# Ensure required env vars are set
check_env("OLLAMA_ENDPOINT")
# Configure structlog
structlog.configure(
processors=[
structlog.stdlib.filter_by_level,
structlog.stdlib.add_logger_name,
structlog.stdlib.add_log_level,
structlog.processors.TimeStamper(fmt="iso"),
structlog.processors.JSONRenderer(),
],
logger_factory=structlog.stdlib.LoggerFactory(),
wrapper_class=structlog.stdlib.BoundLogger,
cache_logger_on_first_use=True,
)
# Set up basic logging configuration
logging.basicConfig(level=logging.ERROR)
# Connect to MLFlow
mlflow.set_tracking_uri("http://10.0.0.2:5000")
experiment = mlflow.get_experiment_by_name("visual-semiotic-ai-analysis")
experiment_id = experiment.experiment_id if experiment else None
# Calculate accuracy
for category, initial_prompt in CATEGORY_PROMPT_MAP.items():
step = 0
run = mlflow.start_run(
experiment_id=experiment_id,
run_name=RUN_NAME,
tags={
"model": MODEL,
"category": category,
},
)
mlflow.log_param(f"step_{step}_prompt", initial_prompt)
print(f"Calculating accuracy for category: {category}")
accuracy = calculate_model_prompt_accuracy(MODEL, initial_prompt, category)
mlflow.log_metric("accuracy", accuracy, step=step)
# Log to MLFlow
prompt_accuracy_map = {initial_prompt: accuracy}
print(f"Prompt accuracy for category '{category}': {accuracy:.2%}")
# Loop to refine prompt based on accuracy
for i in range(99):
step += 1
try:
# Generate refined prompt
refined_prompt = refine_prompt(MODEL, prompt_accuracy_map)
mlflow.log_param(f"step_{step}_prompt", refined_prompt)
print(f"Refined prompt: {refined_prompt}")
except Exception as e:
print(f"Error refining prompt: {e}")
continue
try:
# Calculate accuracy for refined prompt
accuracy = calculate_model_prompt_accuracy(
MODEL, refined_prompt, category
)
mlflow.log_metric("accuracy", accuracy, step=step)
except Exception as e:
print(f"Error calculating accuracy: {e}")
continue
# Update prompt accuracy map
# prompt_accuracy_map[refined_prompt] = accuracy
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."
)
break
mlflow.end_run()