This commit is contained in:
Ján Pták 2026-08-15 01:00:09 +02:00
parent b17b5ed7a4
commit acbb3e82c9

View File

@ -3,14 +3,14 @@ from __future__ import annotations
import argparse import argparse
import csv import csv
import json import json
import sqlite3
import sys import sys
from collections import defaultdict
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
PROJECT_ROOT = Path(__file__).resolve().parents[1] PROJECT_ROOT = Path(
__file__
).resolve().parents[1]
if str(PROJECT_ROOT) not in sys.path: if str(PROJECT_ROOT) not in sys.path:
sys.path.insert( sys.path.insert(
@ -19,798 +19,22 @@ if str(PROJECT_ROOT) not in sys.path:
) )
from evaluation.metrics import (
aggregate_by_category,
aggregate_by_difficulty,
aggregate_metrics,
evaluate_question_mode,
)
from evaluation.retrieval_runner import (
EVALUATION_MODES,
count_splits,
filter_questions_by_split,
load_index_document_paths,
load_questions,
retrieve_all_modes,
validate_dataset,
)
from scripts.common import DB_FILE from scripts.common import DB_FILE
from scripts.search_utils import (
DEFAULT_CANDIDATE_MULTIPLIER,
MIN_CANDIDATES,
add_fts_metadata,
add_vector_metadata,
build_match_queries,
diversify_results,
fuse_hybrid_results,
run_fts_query,
run_vector_query,
verify_search_schema,
)
EVALUATION_K_VALUES = (
1,
3,
5,
)
EVALUATION_MODES = (
"fts",
"vector",
"hybrid",
)
VALID_SPLITS = (
"dev",
"test",
)
def load_questions(
path: Path,
) -> list[dict[str, Any]]:
if not path.exists():
raise FileNotFoundError(
f"Evaluačný dataset neexistuje: {path}"
)
with path.open(
"r",
encoding="utf-8",
) as file:
data = json.load(
file
)
if not isinstance(
data,
list,
):
raise ValueError(
"questions.json musí obsahovať JSON pole"
)
questions: list[
dict[str, Any]
] = []
seen_ids: set[str] = set()
for index, item in enumerate(
data,
start=1,
):
if not isinstance(
item,
dict,
):
raise ValueError(
"Každá evaluačná otázka musí byť "
f"JSON objekt. Chyba pri položke {index}."
)
question_id = item.get(
"id"
)
question = item.get(
"question"
)
split = item.get(
"split"
)
expected_documents = item.get(
"expected_documents"
)
if (
not isinstance(
question_id,
str,
)
or not question_id.strip()
):
raise ValueError(
f"Položka {index} nemá platné id"
)
if question_id in seen_ids:
raise ValueError(
f"Dataset obsahuje duplicitné id: {question_id}"
)
seen_ids.add(
question_id
)
if (
not isinstance(
question,
str,
)
or not question.strip()
):
raise ValueError(
f"{question_id}: chýba otázka"
)
if split not in VALID_SPLITS:
raise ValueError(
f"{question_id}: split musí byť "
"'dev' alebo 'test'"
)
if (
not isinstance(
expected_documents,
list,
)
or not expected_documents
):
raise ValueError(
f"{question_id}: chýba expected_documents"
)
if not all(
isinstance(
value,
str,
)
and value.strip()
for value in expected_documents
):
raise ValueError(
f"{question_id}: expected_documents "
"musí obsahovať neprázdne reťazce"
)
questions.append(
item
)
return questions
def filter_questions_by_split(
questions: list[
dict[str, Any]
],
split: str,
) -> list[dict[str, Any]]:
if split == "all":
return questions
return [
item
for item in questions
if item.get("split") == split
]
def count_splits(
questions: list[
dict[str, Any]
],
) -> dict[str, int]:
counts = {
"dev": 0,
"test": 0,
}
for item in questions:
split = item.get(
"split"
)
if split in counts:
counts[
split
] += 1
return counts
def load_index_document_paths(
db_file: Path,
) -> set[str]:
if not db_file.exists():
raise FileNotFoundError(
f"Databáza neexistuje: {db_file}"
)
with sqlite3.connect(
db_file,
timeout=5.0,
) as conn:
rows = conn.execute(
"""
SELECT DISTINCT
document_path
FROM chunks
ORDER BY document_path
"""
).fetchall()
return {
str(row[0])
for row in rows
}
def validate_dataset(
questions: list[
dict[str, Any]
],
indexed_documents: set[str],
) -> list[dict[str, str]]:
missing: list[
dict[str, str]
] = []
for item in questions:
question_id = str(
item["id"]
)
for document_path in item[
"expected_documents"
]:
if (
document_path
not in indexed_documents
):
missing.append(
{
"question_id": (
question_id
),
"document_path": (
document_path
),
}
)
return missing
def retrieve_all_modes(
db_file: Path,
query: str,
*,
limit: int,
published_only: bool,
max_per_document: int,
) -> dict[
str,
list[dict[str, Any]],
]:
clean_query = query.strip()
if not clean_query:
return {
mode: []
for mode in EVALUATION_MODES
}
match_queries = build_match_queries(
clean_query
)
candidate_limit = max(
MIN_CANDIDATES,
(
limit
* DEFAULT_CANDIDATE_MULTIPLIER
),
)
with sqlite3.connect(
db_file,
timeout=5.0,
) as conn:
conn.row_factory = (
sqlite3.Row
)
conn.execute(
"PRAGMA query_only = ON"
)
verify_search_schema(
conn
)
# -------------------------
# FTS5
# -------------------------
fts_candidates: list[
dict[str, Any]
] = []
used_strategies: list[
str
] = []
for (
strategy,
match_query,
) in match_queries:
rows = run_fts_query(
conn,
match_query,
candidate_limit,
published_only,
)
if not rows:
continue
for row in rows:
row[
"strategy"
] = strategy
fts_candidates = rows
used_strategies = [
strategy
]
break
fts_results = add_fts_metadata(
conn,
clean_query,
fts_candidates,
)
# -------------------------
# Embeddings
# -------------------------
vector_candidates = (
run_vector_query(
conn,
clean_query,
candidate_limit,
published_only,
)
)
vector_results = (
add_vector_metadata(
conn,
vector_candidates,
)
)
# -------------------------
# Hybrid
# -------------------------
hybrid_vector_results = (
vector_results
)
if (
used_strategies
and used_strategies[0]
in {
"all_terms",
"prefix_terms",
}
):
fts_chunk_ids = {
item["chunk_id"]
for item in fts_results
}
hybrid_vector_results = [
item
for item in vector_results
if item["chunk_id"]
in fts_chunk_ids
]
hybrid_results = (
fuse_hybrid_results(
fts_results,
hybrid_vector_results,
)
)
return {
"fts": diversify_results(
fts_results,
limit,
max_per_document,
),
"vector": diversify_results(
vector_results,
limit,
max_per_document,
),
"hybrid": diversify_results(
hybrid_results,
limit,
max_per_document,
),
}
def unique_document_ranking(
results: list[
dict[str, Any]
],
) -> list[dict[str, Any]]:
selected: list[
dict[str, Any]
] = []
seen: set[str] = set()
for item in results:
document_path = str(
item["document_path"]
)
if document_path in seen:
continue
seen.add(
document_path
)
selected.append(
item
)
return selected
def first_relevant_rank(
ranked_documents: list[
dict[str, Any]
],
expected_documents: set[str],
) -> int | None:
for rank, item in enumerate(
ranked_documents,
start=1,
):
if (
item["document_path"]
in expected_documents
):
return rank
return None
def recall_at_k(
ranked_documents: list[
dict[str, Any]
],
expected_documents: set[str],
k: int,
) -> float:
if not expected_documents:
return 0.0
retrieved = {
str(
item["document_path"]
)
for item in ranked_documents[
:k
]
}
relevant_retrieved = (
retrieved
& expected_documents
)
return (
len(
relevant_retrieved
)
/ len(
expected_documents
)
)
def evaluate_question_mode(
question: dict[str, Any],
mode: str,
results: list[
dict[str, Any]
],
) -> dict[str, Any]:
ranked_documents = (
unique_document_ranking(
results
)
)
expected_documents = {
str(value)
for value in question[
"expected_documents"
]
}
rank = first_relevant_rank(
ranked_documents,
expected_documents,
)
reciprocal_rank = (
0.0
if rank is None
else 1.0 / rank
)
row: dict[str, Any] = {
"id": question[
"id"
],
"split": question.get(
"split"
),
"category": question.get(
"category",
"unknown",
),
"difficulty": question.get(
"difficulty",
"unknown",
),
"question": question[
"question"
],
"mode": mode,
"expected_documents": sorted(
expected_documents
),
"first_relevant_rank": rank,
"reciprocal_rank": round(
reciprocal_rank,
6,
),
"top_documents": [
item[
"document_path"
]
for item in ranked_documents
],
"top_source_urls": [
item.get(
"source_url"
)
for item in ranked_documents
],
}
for k in EVALUATION_K_VALUES:
row[
f"hit_at_{k}"
] = (
1
if (
rank is not None
and rank <= k
)
else 0
)
row[
f"recall_at_{k}"
] = round(
recall_at_k(
ranked_documents,
expected_documents,
k,
),
6,
)
return row
def average(
values: list[
float
],
) -> float:
if not values:
return 0.0
return (
sum(values)
/ len(values)
)
def aggregate_metrics(
rows: list[
dict[str, Any]
],
) -> dict[str, Any]:
if not rows:
return {
"questions": 0,
"hit_at_1": 0.0,
"hit_at_3": 0.0,
"hit_at_5": 0.0,
"mrr": 0.0,
"recall_at_5": 0.0,
}
return {
"questions": len(
rows
),
"hit_at_1": round(
average(
[
float(
row[
"hit_at_1"
]
)
for row in rows
]
),
6,
),
"hit_at_3": round(
average(
[
float(
row[
"hit_at_3"
]
)
for row in rows
]
),
6,
),
"hit_at_5": round(
average(
[
float(
row[
"hit_at_5"
]
)
for row in rows
]
),
6,
),
"mrr": round(
average(
[
float(
row[
"reciprocal_rank"
]
)
for row in rows
]
),
6,
),
"recall_at_5": round(
average(
[
float(
row[
"recall_at_5"
]
)
for row in rows
]
),
6,
),
}
def aggregate_by_category(
rows: list[
dict[str, Any]
],
) -> dict[
str,
dict[str, Any],
]:
grouped: dict[
str,
list[dict[str, Any]],
] = defaultdict(
list
)
for row in rows:
category = str(
row.get(
"category",
"unknown",
)
)
grouped[
category
].append(
row
)
return {
category: aggregate_metrics(
category_rows
)
for (
category,
category_rows,
) in sorted(
grouped.items()
)
}
def aggregate_by_difficulty(
rows: list[
dict[str, Any]
],
) -> dict[
str,
dict[str, Any],
]:
grouped: dict[
str,
list[dict[str, Any]],
] = defaultdict(
list
)
for row in rows:
difficulty = str(
row.get(
"difficulty",
"unknown",
)
)
grouped[
difficulty
].append(
row
)
return {
difficulty: aggregate_metrics(
difficulty_rows
)
for (
difficulty,
difficulty_rows,
) in sorted(
grouped.items()
)
}
def print_summary( def print_summary(