Fix db path
This commit is contained in:
parent
2855dd7c13
commit
ca45f8e5d4
2 changed files with 18 additions and 7 deletions
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue