Compare commits
3
Commits
6075848c91
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5faf3ef2a7 | ||
|
|
007161a927 | ||
|
|
d804cc4a13 |
@@ -83,4 +83,5 @@ dev = [
|
|||||||
"ruff>=0.13.0",
|
"ruff>=0.13.0",
|
||||||
"safety>=3.2.11",
|
"safety>=3.2.11",
|
||||||
"testcontainers>=4.13.0",
|
"testcontainers>=4.13.0",
|
||||||
|
"types-requests>=2.32.4.20250913",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -24,7 +24,8 @@ from data_store.dto import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
RUN_NAME = f"query-ai-model_{datetime.now().strftime('%Y%m%d-%H%M%S')}"
|
RUN_NAME = f"query-ai-model_{datetime.now().strftime('%Y%m%d-%H%M%S')}"
|
||||||
MODEL = "llava:13b"
|
IMAGE_MODEL = "llava:13b"
|
||||||
|
PROMPT_MODEL = "mistral-small:22b-instruct-2409-q4_K_M"
|
||||||
CATEGORY_PROMPT_MAP = {
|
CATEGORY_PROMPT_MAP = {
|
||||||
"angle": (
|
"angle": (
|
||||||
"In the context of visual angle in the Kress and van Leeuwen framework, "
|
"In the context of visual angle in the Kress and van Leeuwen framework, "
|
||||||
@@ -204,6 +205,13 @@ def calculate_model_prompt_accuracy(
|
|||||||
total_score += score
|
total_score += score
|
||||||
# Calculate accuracy
|
# Calculate accuracy
|
||||||
jobs_evaluated = total_jobs - jobs_skipped
|
jobs_evaluated = total_jobs - jobs_skipped
|
||||||
|
mlflow.log_metrics(
|
||||||
|
{
|
||||||
|
"total_jobs": total_jobs,
|
||||||
|
"jobs_skipped": jobs_skipped,
|
||||||
|
"jobs_evaluated": jobs_evaluated,
|
||||||
|
}
|
||||||
|
)
|
||||||
print(
|
print(
|
||||||
f"Total Jobs: {total_jobs}, Jobs Skipped: {jobs_skipped}, Jobs Evaluated: {jobs_evaluated}"
|
f"Total Jobs: {total_jobs}, Jobs Skipped: {jobs_skipped}, Jobs Evaluated: {jobs_evaluated}"
|
||||||
)
|
)
|
||||||
@@ -278,20 +286,20 @@ if __name__ == "__main__":
|
|||||||
experiment = mlflow.get_experiment_by_name("visual-semiotic-ai-analysis")
|
experiment = mlflow.get_experiment_by_name("visual-semiotic-ai-analysis")
|
||||||
experiment_id = experiment.experiment_id if experiment else None
|
experiment_id = experiment.experiment_id if experiment else None
|
||||||
# Calculate accuracy
|
# Calculate accuracy
|
||||||
category_events = {}
|
|
||||||
for category, initial_prompt in CATEGORY_PROMPT_MAP.items():
|
for category, initial_prompt in CATEGORY_PROMPT_MAP.items():
|
||||||
step = 0
|
step = 0
|
||||||
run = mlflow.start_run(
|
run = mlflow.start_run(
|
||||||
experiment_id=experiment_id,
|
experiment_id=experiment_id,
|
||||||
run_name=RUN_NAME,
|
run_name=RUN_NAME,
|
||||||
tags={
|
tags={
|
||||||
"model": MODEL,
|
"image_model": IMAGE_MODEL,
|
||||||
|
"prompt_model": PROMPT_MODEL,
|
||||||
"category": category,
|
"category": category,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
mlflow.log_param(f"step_{step}_prompt", initial_prompt)
|
mlflow.log_param(f"step_{step}_prompt", initial_prompt)
|
||||||
print(f"Calculating accuracy for category: {category}")
|
print(f"Calculating accuracy for category: {category}")
|
||||||
accuracy = calculate_model_prompt_accuracy(MODEL, initial_prompt, category)
|
accuracy = calculate_model_prompt_accuracy(IMAGE_MODEL, initial_prompt, category)
|
||||||
mlflow.log_metric("accuracy", accuracy, step=step)
|
mlflow.log_metric("accuracy", accuracy, step=step)
|
||||||
# Log to MLFlow
|
# Log to MLFlow
|
||||||
prompt_accuracy_map = {initial_prompt: accuracy}
|
prompt_accuracy_map = {initial_prompt: accuracy}
|
||||||
@@ -301,7 +309,7 @@ if __name__ == "__main__":
|
|||||||
step += 1
|
step += 1
|
||||||
try:
|
try:
|
||||||
# Generate refined prompt
|
# Generate refined prompt
|
||||||
refined_prompt = refine_prompt(MODEL, prompt_accuracy_map)
|
refined_prompt = refine_prompt(PROMPT_MODEL, prompt_accuracy_map)
|
||||||
mlflow.log_param(f"step_{step}_prompt", refined_prompt)
|
mlflow.log_param(f"step_{step}_prompt", refined_prompt)
|
||||||
print(f"Refined prompt: {refined_prompt}")
|
print(f"Refined prompt: {refined_prompt}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -309,8 +317,10 @@ if __name__ == "__main__":
|
|||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
# Calculate accuracy for refined prompt
|
# Calculate accuracy for refined prompt
|
||||||
accuracy = calculate_model_prompt_accuracy(MODEL, refined_prompt, category)
|
accuracy = calculate_model_prompt_accuracy(
|
||||||
mlflow.log_metric(f"accuracy", accuracy, step=step)
|
IMAGE_MODEL, refined_prompt, category
|
||||||
|
)
|
||||||
|
mlflow.log_metric("accuracy", accuracy, step=step)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"Error calculating accuracy: {e}")
|
print(f"Error calculating accuracy: {e}")
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -2629,6 +2629,18 @@ wheels = [
|
|||||||
{ url = "http://10.0.0.2:5001/index/types-pytz/types_pytz-2025.2.0.20250809-py3-none-any.whl", hash = "sha256:4f55ed1b43e925cf851a756fe1707e0f5deeb1976e15bf844bcaa025e8fbd0db" },
|
{ url = "http://10.0.0.2:5001/index/types-pytz/types_pytz-2025.2.0.20250809-py3-none-any.whl", hash = "sha256:4f55ed1b43e925cf851a756fe1707e0f5deeb1976e15bf844bcaa025e8fbd0db" },
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "types-requests"
|
||||||
|
version = "2.32.4.20250913"
|
||||||
|
source = { registry = "http://10.0.0.2:5001/index/" }
|
||||||
|
dependencies = [
|
||||||
|
{ name = "urllib3" },
|
||||||
|
]
|
||||||
|
sdist = { url = "http://10.0.0.2:5001/index/types-requests/types_requests-2.32.4.20250913.tar.gz", hash = "sha256:abd6d4f9ce3a9383f269775a9835a4c24e5cd6b9f647d64f88aa4613c33def5d" }
|
||||||
|
wheels = [
|
||||||
|
{ url = "http://10.0.0.2:5001/index/types-requests/types_requests-2.32.4.20250913-py3-none-any.whl", hash = "sha256:78c9c1fffebbe0fa487a418e0fa5252017e9c60d1a2da394077f1780f655d7e1" },
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "typing-extensions"
|
name = "typing-extensions"
|
||||||
version = "4.15.0"
|
version = "4.15.0"
|
||||||
@@ -2720,6 +2732,7 @@ dev = [
|
|||||||
{ name = "ruff" },
|
{ name = "ruff" },
|
||||||
{ name = "safety" },
|
{ name = "safety" },
|
||||||
{ name = "testcontainers" },
|
{ name = "testcontainers" },
|
||||||
|
{ name = "types-requests" },
|
||||||
]
|
]
|
||||||
|
|
||||||
[package.metadata]
|
[package.metadata]
|
||||||
@@ -2744,6 +2757,7 @@ dev = [
|
|||||||
{ name = "ruff", specifier = ">=0.13.0" },
|
{ name = "ruff", specifier = ">=0.13.0" },
|
||||||
{ name = "safety", specifier = ">=3.2.11" },
|
{ name = "safety", specifier = ">=3.2.11" },
|
||||||
{ name = "testcontainers", specifier = ">=4.13.0" },
|
{ name = "testcontainers", specifier = ">=4.13.0" },
|
||||||
|
{ name = "types-requests", specifier = ">=2.32.4.20250913" },
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
|
|||||||
Reference in New Issue
Block a user