Fix db path

This commit is contained in:
Yiorgis Gozadinos 2025-06-25 19:01:57 +03:00
parent 2855dd7c13
commit ca45f8e5d4
No known key found for this signature in database
2 changed files with 18 additions and 7 deletions

View file

@ -4,6 +4,8 @@ from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config from haiku.rag.config import Config
from haiku.rag.qa.base import QABase from haiku.rag.qa.base import QABase
OLLAMA_OPTIONS = {"temperature": 0.0, "seed": 42, "num_ctx": 64000}
class QA(QABase): class QA(QABase):
def __init__(self, client: HaikuRAG, model: str = Config.QA_MODEL): def __init__(self, client: HaikuRAG, model: str = Config.QA_MODEL):
@ -48,7 +50,7 @@ class QA(QABase):
model=self._model, model=self._model,
messages=messages, messages=messages,
tools=tools, tools=tools,
options={"temperature": 0.0, "seed": 42}, options=OLLAMA_OPTIONS,
think=False, think=False,
) )
@ -82,7 +84,7 @@ class QA(QABase):
model=self._model, model=self._model,
messages=messages, messages=messages,
think=False, think=False,
options={"temperature": 0.0, "seed": 42}, options=OLLAMA_OPTIONS,
) )
return final_response["message"]["content"] return final_response["message"]["content"]
else: else:

View file

@ -8,16 +8,18 @@ from tqdm import tqdm
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.qa.ollama import QA from haiku.rag.qa.ollama import QA
db_path = Path(__file__).parent / "data" / "benchmark.sqlite"
async def populate_db(): async def populate_db():
if (Path(__file__).parent / "data" / "benchmark.sqlite").exists(): if (db_path).exists():
print("Benchmark database already exists. Skipping creation.") print("Benchmark database already exists. Skipping creation.")
return return
ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore
corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories") corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories")
async with HaikuRAG(Path(__file__).parent / "benchmark.sqlite") as rag: async with HaikuRAG(db_path) as rag:
for i, doc in enumerate(tqdm(corpus)): for i, doc in enumerate(tqdm(corpus)):
await rag.create_document( await rag.create_document(
content=doc["document_extracted"], # type: ignore content=doc["document_extracted"], # type: ignore
@ -34,7 +36,7 @@ async def run_match_benchmark():
correct_at_3 = 0 correct_at_3 = 0
total_queries = 0 total_queries = 0
async with HaikuRAG(Path(__file__).parent / "benchmark.sqlite") as rag: async with HaikuRAG(db_path) as rag:
for i, doc in enumerate(tqdm(corpus)): for i, doc in enumerate(tqdm(corpus)):
doc_id = doc["document_id"] # type: ignore doc_id = doc["document_id"] # type: ignore
matches = await rag.search( matches = await rag.search(
@ -73,16 +75,19 @@ async def run_match_benchmark():
return {"recall@1": recall_at_1, "recall@2": recall_at_2, "recall@3": recall_at_3} return {"recall@1": recall_at_1, "recall@2": recall_at_2, "recall@3": recall_at_3}
async def run_qa_benchmark(): async def run_qa_benchmark(k: int | None = None):
"""Run QA benchmarking on the corpus.""" """Run QA benchmarking on the corpus."""
ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore
corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories") corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories")
if k is not None:
corpus = corpus.select(range(min(k, len(corpus))))
judge = LLMJudge() judge = LLMJudge()
correct_answers = 0 correct_answers = 0
total_questions = 0 total_questions = 0
async with HaikuRAG(Path(__file__).parent / "benchmark.sqlite") as rag: async with HaikuRAG(db_path) as rag:
qa = QA(rag) qa = QA(rag)
for i, doc in enumerate(tqdm(corpus, desc="QA Benchmarking")): for i, doc in enumerate(tqdm(corpus, desc="QA Benchmarking")):
@ -93,6 +98,10 @@ async def run_qa_benchmark():
is_equivalent = await judge.judge_answers( is_equivalent = await judge.judge_answers(
question, generated_answer, expected_answer question, generated_answer, expected_answer
) )
print(f"Question: {question}")
print(f"Expected: {expected_answer}")
print(f"Generated: {generated_answer}")
print(f"Equivalent: {is_equivalent}\n")
if is_equivalent: if is_equivalent:
correct_answers += 1 correct_answers += 1