diff --git a/evaluation/retrieval_runner.py b/evaluation/retrieval_runner.py new file mode 100644 index 0000000..8a08353 --- /dev/null +++ b/evaluation/retrieval_runner.py @@ -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, + ), + }