rozdelenie_kodu
This commit is contained in:
parent
ba24b9b6f2
commit
1b3e4f79c8
426
evaluation/retrieval_runner.py
Normal file
426
evaluation/retrieval_runner.py
Normal file
@ -0,0 +1,426 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import sqlite3
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
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_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,
|
||||||
|
),
|
||||||
|
}
|
||||||
Loading…
Reference in New Issue
Block a user