Adapt how we measure recall when using datasets with multiple sources

This commit is contained in:
Yiorgis Gozadinos 2025-09-29 20:02:37 +03:00
parent eb2fc67e1c
commit 842e166041
No known key found for this signature in database
2 changed files with 29 additions and 32 deletions

View file

@ -17,13 +17,10 @@ The recall obtained is ~0.79 for matching in the top result, raising to ~0.91 fo
| Embedding Model | Document in top 1 | Document in top 3 | Reranker |
|---------------------------------------|-------------------|-------------------|------------------------|
| Ollama / `qwen3-embedding` | 0.81 | 0.95 | None |
| Ollama / `mxbai-embed-large` | 0.79 | 0.91 | None |
| Ollama / `mxbai-embed-large` | 0.90 | 0.95 | `mxbai-rerank-base-v2` |
| Ollama / `nomic-embed-text-v1.5` | 0.74 | 0.90 | None |
| Ollama / `qwen3-embedding` | 0.81 | 0.95 | None |
<!-- | OpenAI / `text-embeddings-3-small` | 0.75 | 0.88 | None |
| OpenAI / `text-embeddings-3-small` | 0.75 | 0.88 | None |
| OpenAI / `text-embeddings-3-small` | 0.83 | 0.90 | Cohere / `rerank-v3.5` | -->
## Question/Answer evaluation
@ -39,5 +36,3 @@ determine whether the answer is correct. The obtained accuracy is as follows:
| Ollama / `mxbai-embed-large` | Ollama / `qwen3:0.6b` | 0.28 | None |
Note the significant degradation when very small models are used such as `qwen3:0.6b`.
<!-- | Ollama / `mxbai-embed-large` | Anthropic / `Claude Sonnet 3.7` | 0.79 | None |
| OpenAI / `text-embeddings-3-small` | OpenAI / `gpt-4-turbo` | 0.62 | None | -->

View file

@ -75,9 +75,11 @@ async def run_retrieval_benchmark(spec: DatasetSpec) -> dict[str, float] | None:
corpus = spec.retrieval_loader()
correct_at_1 = 0
correct_at_2 = 0
correct_at_3 = 0
recall_totals = {
1: 0.0,
3: 0.0,
5: 0.0,
}
total_queries = 0
with Progress() as progress:
@ -92,30 +94,30 @@ async def run_retrieval_benchmark(spec: DatasetSpec) -> dict[str, float] | None:
progress.advance(task)
continue
matches = await rag.search(query=sample.question, limit=3)
matches = await rag.search(query=sample.question, limit=5)
if not matches:
progress.advance(task)
continue
total_queries += 1
for position, (chunk, _) in enumerate(matches):
retrieved = (
await rag.get_document_by_id(chunk.document_id)
if chunk.document_id is not None
else None
)
if retrieved and _is_relevant_match(retrieved.uri, sample):
if position == 0:
correct_at_1 += 1
correct_at_2 += 1
correct_at_3 += 1
elif position == 1:
correct_at_2 += 1
correct_at_3 += 1
elif position == 2:
correct_at_3 += 1
break
retrieved_uris: list[str] = []
for chunk, _ in matches:
if chunk.document_id is None:
continue
retrieved_doc = await rag.get_document_by_id(chunk.document_id)
if retrieved_doc and retrieved_doc.uri:
retrieved_uris.append(retrieved_doc.uri)
# Compute per-query recall@K by counting how many relevant
# documents are retrieved within the first K results and
# averaging these fractions across all queries.
for cutoff in (1, 3, 5):
top_k = set(retrieved_uris[:cutoff])
relevant = set(sample.expected_uris)
if relevant:
matched = len(top_k & relevant)
recall_totals[cutoff] += matched / len(relevant)
progress.advance(task)
@ -123,20 +125,20 @@ async def run_retrieval_benchmark(spec: DatasetSpec) -> dict[str, float] | None:
console.print("No retrieval cases to evaluate.")
return None
recall_at_1 = correct_at_1 / total_queries
recall_at_2 = correct_at_2 / total_queries
recall_at_3 = correct_at_3 / total_queries
recall_at_1 = recall_totals[1] / total_queries
recall_at_3 = recall_totals[3] / total_queries
recall_at_5 = recall_totals[5] / total_queries
console.print("\n=== Retrieval Benchmark Results ===", style="bold cyan")
console.print(f"Total queries: {total_queries}")
console.print(f"Recall@1: {recall_at_1:.4f}")
console.print(f"Recall@2: {recall_at_2:.4f}")
console.print(f"Recall@3: {recall_at_3:.4f}")
console.print(f"Recall@5: {recall_at_5:.4f}")
return {
"recall@1": recall_at_1,
"recall@2": recall_at_2,
"recall@3": recall_at_3,
"recall@5": recall_at_5,
}