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