Spaces:
Running
Running
GitHub Actions commited on
Commit ·
eeed2cb
1
Parent(s): 7667ee3
Sync from GitHub Actions
Browse files
visual_product_search/evaluation/__init__.py
CHANGED
|
@@ -1,3 +1,4 @@
|
|
|
|
|
| 1 |
from visual_product_search.evaluation.metrics import (
|
| 2 |
average_metric_dicts,
|
| 3 |
average_precision_at_k,
|
|
@@ -10,6 +11,7 @@ from visual_product_search.evaluation.metrics import (
|
|
| 10 |
)
|
| 11 |
|
| 12 |
__all__ = [
|
|
|
|
| 13 |
"average_metric_dicts",
|
| 14 |
"average_precision_at_k",
|
| 15 |
"dcg_at_k",
|
|
|
|
| 1 |
+
from visual_product_search.evaluation.evaluator import ImageToImageEvaluator
|
| 2 |
from visual_product_search.evaluation.metrics import (
|
| 3 |
average_metric_dicts,
|
| 4 |
average_precision_at_k,
|
|
|
|
| 11 |
)
|
| 12 |
|
| 13 |
__all__ = [
|
| 14 |
+
"ImageToImageEvaluator",
|
| 15 |
"average_metric_dicts",
|
| 16 |
"average_precision_at_k",
|
| 17 |
"dcg_at_k",
|
visual_product_search/evaluation/error_analysis.py
CHANGED
|
@@ -0,0 +1,50 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import pandas as pd
|
| 2 |
+
|
| 3 |
+
def normalize_value(value):
|
| 4 |
+
if pd.isna(value):
|
| 5 |
+
return ""
|
| 6 |
+
return str(value).strip().lower()
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def row_matches(query_row, candidate_row, fields):
|
| 10 |
+
for field in fields:
|
| 11 |
+
if field not in query_row.index or field not in candidate_row.index:
|
| 12 |
+
return False
|
| 13 |
+
if normalize_value(query_row[field]) != normalize_value(candidate_row[field]):
|
| 14 |
+
return False
|
| 15 |
+
|
| 16 |
+
return True
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def detect_error_type(query_row, candidate_row):
|
| 20 |
+
checks = [
|
| 21 |
+
("articleType", "wrong_articleType"),
|
| 22 |
+
("subCategory", "wrong_subCategory"),
|
| 23 |
+
("gender", "wrong_gender"),
|
| 24 |
+
("masterCategory", "wrong_masterCategory"),
|
| 25 |
+
]
|
| 26 |
+
|
| 27 |
+
for field, error_name in checks:
|
| 28 |
+
if field in query_row.index and field in candidate_row.index:
|
| 29 |
+
if normalize_value(query_row[field]) != normalize_value(candidate_row[field]):
|
| 30 |
+
return error_name
|
| 31 |
+
return "ranking_error"
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def build_error_case(query_row, candidate_row, score, query_index, candidate_index):
|
| 35 |
+
return {
|
| 36 |
+
"query_index": int(query_index),
|
| 37 |
+
"top1_index": int(candidate_index),
|
| 38 |
+
"query_image": str(query_row.get("image_path", query_row.get("filename", query_row.get("image", "")))),
|
| 39 |
+
"query_articleType": str(query_row.get("articleType", "")),
|
| 40 |
+
"query_subCategory": str(query_row.get("subCategory", "")),
|
| 41 |
+
"query_gender": str(query_row.get("gender", "")),
|
| 42 |
+
"query_masterCategory": str(query_row.get("masterCategory", "")),
|
| 43 |
+
"top1_image": str(candidate_row.get("image_path", candidate_row.get("filename", candidate_row.get("image", "")))),
|
| 44 |
+
"top1_articleType": str(candidate_row.get("articleType", "")),
|
| 45 |
+
"top1_subCategory": str(candidate_row.get("subCategory", "")),
|
| 46 |
+
"top1_gender": str(candidate_row.get("gender", "")),
|
| 47 |
+
"top1_masterCategory": str(candidate_row.get("masterCategory", "")),
|
| 48 |
+
"similarity_score": float(score),
|
| 49 |
+
"error_type": detect_error_type(query_row, candidate_row),
|
| 50 |
+
}
|
visual_product_search/evaluation/evaluate_image_to_image.py
CHANGED
|
@@ -0,0 +1,17 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import sys
|
| 3 |
+
from visual_product_search.evaluation.evaluator import ImageToImageEvaluator
|
| 4 |
+
from visual_product_search.exception import ExceptionHandle
|
| 5 |
+
|
| 6 |
+
def main():
|
| 7 |
+
try:
|
| 8 |
+
evaluator = ImageToImageEvaluator()
|
| 9 |
+
result = evaluator.run()
|
| 10 |
+
print(json.dumps(result, indent=2))
|
| 11 |
+
|
| 12 |
+
except Exception as e:
|
| 13 |
+
raise ExceptionHandle(e, sys)
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
if __name__ == "__main__":
|
| 17 |
+
main()
|
visual_product_search/evaluation/evaluator.py
ADDED
|
@@ -0,0 +1,205 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from collections import defaultdict
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
import numpy as np
|
| 4 |
+
from visual_product_search.evaluation.data_loader import EvaluationDataLoader
|
| 5 |
+
from visual_product_search.evaluation.embedding_store import EmbeddingStore
|
| 6 |
+
from visual_product_search.evaluation.error_analysis import build_error_case
|
| 7 |
+
from visual_product_search.evaluation.metrics import average_metric_dicts, evaluate_ranking
|
| 8 |
+
from visual_product_search.evaluation.relevance import RelevanceComputer
|
| 9 |
+
from visual_product_search.evaluation.report import EvaluationReportWriter
|
| 10 |
+
from visual_product_search.evaluation.similarity import SimilaritySearcher
|
| 11 |
+
from visual_product_search.logger import logging
|
| 12 |
+
from visual_product_search.utils.config import load_config
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class ImageToImageEvaluator:
|
| 16 |
+
def __init__(self, config_path="config/model.yaml"):
|
| 17 |
+
self.config_path = config_path
|
| 18 |
+
self.config = load_config(config_path)
|
| 19 |
+
self.root = Path.cwd()
|
| 20 |
+
|
| 21 |
+
self.model_config = self.config.get("model", {})
|
| 22 |
+
self.data_config = self.config.get("data", {})
|
| 23 |
+
self.evaluation_config = self.config.get("evaluation", {})
|
| 24 |
+
|
| 25 |
+
self.k_values = [int(k) for k in self.evaluation_config.get("k_values", [1, 5, 10])]
|
| 26 |
+
self.max_k = max(self.k_values)
|
| 27 |
+
|
| 28 |
+
self.query_batch_size = int(self.evaluation_config.get("query_batch_size", 64))
|
| 29 |
+
self.max_queries = self.evaluation_config.get("max_queries")
|
| 30 |
+
self.seed = int(self.evaluation_config.get("seed", 42))
|
| 31 |
+
self.exclude_self_match = bool(self.evaluation_config.get("exclude_self_match", True))
|
| 32 |
+
|
| 33 |
+
self.data_loader = EvaluationDataLoader(self.config, self.root)
|
| 34 |
+
self.embedding_store = EmbeddingStore(self.config, self.root)
|
| 35 |
+
self.report_writer = EvaluationReportWriter(self.config, self.root)
|
| 36 |
+
|
| 37 |
+
def select_query_indices(self, total_items):
|
| 38 |
+
indices = np.arange(total_items)
|
| 39 |
+
if self.max_queries is None:
|
| 40 |
+
return indices
|
| 41 |
+
|
| 42 |
+
max_queries = int(self.max_queries)
|
| 43 |
+
if max_queries <= 0 or max_queries >= total_items:
|
| 44 |
+
return indices
|
| 45 |
+
|
| 46 |
+
rng = np.random.default_rng(self.seed)
|
| 47 |
+
selected = rng.choice(indices, size=max_queries, replace=False)
|
| 48 |
+
|
| 49 |
+
return np.sort(selected)
|
| 50 |
+
|
| 51 |
+
def get_category_field(self, metadata):
|
| 52 |
+
if "articleType" in metadata.columns:
|
| 53 |
+
return "articleType"
|
| 54 |
+
|
| 55 |
+
if "subCategory" in metadata.columns:
|
| 56 |
+
return "subCategory"
|
| 57 |
+
|
| 58 |
+
return None
|
| 59 |
+
|
| 60 |
+
def build_metrics_output(self, metadata, embeddings, query_indices, relevance_computer, metric_storage):
|
| 61 |
+
return {
|
| 62 |
+
"evaluation_type": "image_to_image_retrieval",
|
| 63 |
+
"model": (
|
| 64 |
+
self.model_config.get("repo_id")
|
| 65 |
+
or self.model_config.get("new_model")
|
| 66 |
+
or self.model_config.get("name")
|
| 67 |
+
),
|
| 68 |
+
"dataset": self.data_config.get(
|
| 69 |
+
"dataset_name",
|
| 70 |
+
"paramaggarwal/fashion-product-images-dataset",
|
| 71 |
+
),
|
| 72 |
+
"num_queries": int(len(query_indices)),
|
| 73 |
+
"num_indexed_images": int(len(metadata)),
|
| 74 |
+
"embedding_dimension": int(embeddings.shape[1]),
|
| 75 |
+
"k_values": self.k_values,
|
| 76 |
+
"exclude_self_match": self.exclude_self_match,
|
| 77 |
+
"relevance_definitions": relevance_computer.levels,
|
| 78 |
+
"metrics": {
|
| 79 |
+
level: average_metric_dicts(values)
|
| 80 |
+
for level, values in metric_storage.items()
|
| 81 |
+
},
|
| 82 |
+
}
|
| 83 |
+
|
| 84 |
+
def build_category_rows(self, category_storage):
|
| 85 |
+
category_rows = []
|
| 86 |
+
for level_name, categories in category_storage.items():
|
| 87 |
+
for category, values in categories.items():
|
| 88 |
+
row = {
|
| 89 |
+
"relevance_level": level_name,
|
| 90 |
+
"category": category,
|
| 91 |
+
"num_queries": len(values),
|
| 92 |
+
}
|
| 93 |
+
|
| 94 |
+
row.update(average_metric_dicts(values))
|
| 95 |
+
category_rows.append(row)
|
| 96 |
+
|
| 97 |
+
return category_rows
|
| 98 |
+
|
| 99 |
+
def run(self):
|
| 100 |
+
metadata = self.data_loader.load()
|
| 101 |
+
embeddings, metadata = self.embedding_store.load_or_build(metadata)
|
| 102 |
+
|
| 103 |
+
if len(embeddings) != len(metadata):
|
| 104 |
+
raise ValueError("Embeddings and metadata must have the same length")
|
| 105 |
+
|
| 106 |
+
query_indices = self.select_query_indices(len(metadata))
|
| 107 |
+
|
| 108 |
+
relevance_computer = RelevanceComputer(metadata, self.evaluation_config)
|
| 109 |
+
searcher = SimilaritySearcher(
|
| 110 |
+
top_k=self.max_k,
|
| 111 |
+
exclude_self_match=self.exclude_self_match,
|
| 112 |
+
)
|
| 113 |
+
metric_storage = {
|
| 114 |
+
level: []
|
| 115 |
+
for level in relevance_computer.levels
|
| 116 |
+
}
|
| 117 |
+
|
| 118 |
+
category_storage = defaultdict(lambda: defaultdict(list))
|
| 119 |
+
error_cases = []
|
| 120 |
+
|
| 121 |
+
category_field = self.get_category_field(metadata)
|
| 122 |
+
for start in range(0, len(query_indices), self.query_batch_size):
|
| 123 |
+
batch_indices = query_indices[start:start + self.query_batch_size]
|
| 124 |
+
query_embeddings = embeddings[batch_indices]
|
| 125 |
+
|
| 126 |
+
retrieved_indices_batch, retrieved_scores_batch = searcher.search_batch(
|
| 127 |
+
query_embeddings=query_embeddings,
|
| 128 |
+
all_embeddings=embeddings,
|
| 129 |
+
query_indices=batch_indices,
|
| 130 |
+
)
|
| 131 |
+
|
| 132 |
+
for row_number, query_index in enumerate(batch_indices):
|
| 133 |
+
query_row = metadata.iloc[query_index]
|
| 134 |
+
retrieved_indices = retrieved_indices_batch[row_number]
|
| 135 |
+
retrieved_scores = retrieved_scores_batch[row_number]
|
| 136 |
+
|
| 137 |
+
if category_field:
|
| 138 |
+
category = str(query_row.get(category_field, "unknown"))
|
| 139 |
+
else:
|
| 140 |
+
category = "unknown"
|
| 141 |
+
|
| 142 |
+
for level_name, fields in relevance_computer.levels.items():
|
| 143 |
+
relevance, total_relevant, _ = relevance_computer.get_relevance_list(
|
| 144 |
+
retrieved_indices=retrieved_indices,
|
| 145 |
+
fields=fields,
|
| 146 |
+
query_index=query_index,
|
| 147 |
+
)
|
| 148 |
+
|
| 149 |
+
metrics = evaluate_ranking(
|
| 150 |
+
relevance=relevance,
|
| 151 |
+
total_relevant=total_relevant,
|
| 152 |
+
k_values=self.k_values,
|
| 153 |
+
)
|
| 154 |
+
|
| 155 |
+
metric_storage[level_name].append(metrics)
|
| 156 |
+
category_storage[level_name][category].append(metrics)
|
| 157 |
+
|
| 158 |
+
strict_fields = relevance_computer.levels.get("strict", ["articleType"])
|
| 159 |
+
|
| 160 |
+
if retrieved_indices:
|
| 161 |
+
_, _, strict_mask = relevance_computer.get_relevance_list(
|
| 162 |
+
retrieved_indices=retrieved_indices,
|
| 163 |
+
fields=strict_fields,
|
| 164 |
+
query_index=query_index,
|
| 165 |
+
)
|
| 166 |
+
|
| 167 |
+
top1_index = retrieved_indices[0]
|
| 168 |
+
top1_score = retrieved_scores[0]
|
| 169 |
+
|
| 170 |
+
if not bool(strict_mask[top1_index]):
|
| 171 |
+
error_cases.append(
|
| 172 |
+
build_error_case(
|
| 173 |
+
query_row=query_row,
|
| 174 |
+
candidate_row=metadata.iloc[top1_index],
|
| 175 |
+
score=top1_score,
|
| 176 |
+
query_index=query_index,
|
| 177 |
+
candidate_index=top1_index,
|
| 178 |
+
)
|
| 179 |
+
)
|
| 180 |
+
|
| 181 |
+
logging.info(
|
| 182 |
+
f"Evaluated {min(start + self.query_batch_size, len(query_indices))}/{len(query_indices)} queries"
|
| 183 |
+
)
|
| 184 |
+
|
| 185 |
+
metrics_output = self.build_metrics_output(
|
| 186 |
+
metadata=metadata,
|
| 187 |
+
embeddings=embeddings,
|
| 188 |
+
query_indices=query_indices,
|
| 189 |
+
relevance_computer=relevance_computer,
|
| 190 |
+
metric_storage=metric_storage,
|
| 191 |
+
)
|
| 192 |
+
|
| 193 |
+
category_rows = self.build_category_rows(category_storage)
|
| 194 |
+
|
| 195 |
+
saved_paths = self.report_writer.save(
|
| 196 |
+
metrics_output=metrics_output,
|
| 197 |
+
category_rows=category_rows,
|
| 198 |
+
error_cases=error_cases,
|
| 199 |
+
)
|
| 200 |
+
|
| 201 |
+
logging.info(f"Evaluation metrics saved to {saved_paths['metrics_path']}")
|
| 202 |
+
logging.info(f"Category breakdown saved to {saved_paths['category_breakdown_path']}")
|
| 203 |
+
logging.info(f"Error cases saved to {saved_paths['error_cases_path']}")
|
| 204 |
+
|
| 205 |
+
return metrics_output
|