added use of larger model for refining prompt
This commit is contained in:
@@ -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}"
|
||||||
)
|
)
|
||||||
@@ -284,13 +292,14 @@ if __name__ == "__main__":
|
|||||||
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}
|
||||||
@@ -300,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,7 +318,7 @@ if __name__ == "__main__":
|
|||||||
try:
|
try:
|
||||||
# Calculate accuracy for refined prompt
|
# Calculate accuracy for refined prompt
|
||||||
accuracy = calculate_model_prompt_accuracy(
|
accuracy = calculate_model_prompt_accuracy(
|
||||||
MODEL, refined_prompt, category
|
IMAGE_MODEL, refined_prompt, category
|
||||||
)
|
)
|
||||||
mlflow.log_metric("accuracy", accuracy, step=step)
|
mlflow.log_metric("accuracy", accuracy, step=step)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
Reference in New Issue
Block a user