Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 16 additions & 7 deletions haystack/components/evaluators/document_map.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,19 +117,28 @@ def run(
for ground_truth, retrieved in zip(ground_truth_documents, retrieved_documents, strict=True):
average_precision = 0.0
average_precision_numerator = 0.0
relevant_documents = 0
retrieved_relevant_documents = 0

ground_truth_values = [val for doc in ground_truth if (val := self._get_comparison_value(doc)) is not None]
# A list keeps the deduplication working for unhashable comparison values, for example when
# document_comparison_field points to a meta key holding a list.
uncredited_ground_truth_values: list[Any] = []
for doc in ground_truth:
value = self._get_comparison_value(doc)
if value is not None and value not in uncredited_ground_truth_values:
uncredited_ground_truth_values.append(value)

total_relevant_documents = len(uncredited_ground_truth_values)
for rank, retrieved_document in enumerate(retrieved):
retrieved_value = self._get_comparison_value(retrieved_document)
if retrieved_value is None:
continue

if retrieved_value in ground_truth_values:
relevant_documents += 1
average_precision_numerator += relevant_documents / (rank + 1)
if relevant_documents > 0:
average_precision = average_precision_numerator / relevant_documents
if retrieved_value in uncredited_ground_truth_values:
uncredited_ground_truth_values.remove(retrieved_value)
retrieved_relevant_documents += 1
average_precision_numerator += retrieved_relevant_documents / (rank + 1)
if total_relevant_documents:
average_precision = average_precision_numerator / total_relevant_documents
individual_scores.append(average_precision)

score = sum(individual_scores) / len(ground_truth_documents)
Expand Down
10 changes: 10 additions & 0 deletions releasenotes/notes/fix-document-map-average-precision.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
---
fixes:
- |
Fix ``DocumentMAPEvaluator`` to include missed relevant documents in the average precision denominator and avoid
crediting duplicate retrievals of the same document.
upgrade:
- |
``DocumentMAPEvaluator`` scores can change because average precision now uses all unique, valid ground-truth
comparison values as its denominator and credits each value at most once. Re-baseline evaluations that relied on
the previous scores.
31 changes: 28 additions & 3 deletions test/components/evaluators/test_document_map.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,18 @@ def test_run_with_nested_meta_comparison():
assert result == {"individual_scores": [1.0, 0.0], "score": 0.5}


def test_run_with_unhashable_meta_comparison():
evaluator = DocumentMAPEvaluator(document_comparison_field="meta.tags")
result = evaluator.run(
ground_truth_documents=[
[Document(content="x", meta={"tags": ["a"]}), Document(content="y", meta={"tags": ["b"]})]
],
retrieved_documents=[[Document(content="z", meta={"tags": ["a"]})]],
)

assert result == {"individual_scores": [0.5], "score": 0.5}


def test_run_with_all_matching():
evaluator = DocumentMAPEvaluator()
result = evaluator.run(
Expand Down Expand Up @@ -95,6 +107,19 @@ def test_run_with_partial_matching():
assert result == {"individual_scores": [1.0, 0.0], "score": 0.5}


@pytest.mark.parametrize(
"retrieved_documents", [[Document(content="A")], [Document(content="A"), Document(content="A")]]
)
def test_run_with_missed_and_duplicate_relevant_documents(retrieved_documents):
evaluator = DocumentMAPEvaluator()
result = evaluator.run(
ground_truth_documents=[[Document(content="A"), Document(content="B")]],
retrieved_documents=[retrieved_documents],
)

assert result == {"individual_scores": [0.5], "score": 0.5}


def test_run_with_complex_data():
evaluator = DocumentMAPEvaluator()
result = evaluator.run(
Expand Down Expand Up @@ -124,12 +149,12 @@ def test_run_with_complex_data():
"individual_scores": [
1.0,
pytest.approx(0.8333333333333333),
1.0,
0.5, # Only one of two relevant documents was retrieved.
pytest.approx(0.5833333333333333),
0.0,
pytest.approx(0.8055555555555555),
pytest.approx(0.8333333333333333), # The duplicate retrieval is not credited again.
],
"score": pytest.approx(0.7037037037037037),
"score": pytest.approx(0.625),
}


Expand Down
Loading