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