Skip to content

Commit f5c63f9

Browse files
authored
Merge pull request #21 from erwallace/evaluation
Evaluation
2 parents edc7898 + e12d5fe commit f5c63f9

11 files changed

Lines changed: 221 additions & 22 deletions

File tree

.gitattributes

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
vectorstore/* filter=lfs diff=lfs merge=lfs -text
2+
vectorstore/*.sqlite filter=lfs diff=lfs merge=lfs -text
3+
vectorstore/**/*.bin filter=lfs diff=lfs merge=lfs -text

evaluation/evaluation.py

Lines changed: 125 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,125 @@
1+
# This script evaluates a set of question-answer pairs using the RAG chatbot.
2+
import json
3+
4+
from loguru import logger
5+
from paper_query.chatbots import HybridQueryChatbot
6+
from paper_query.constants import METRICS_JSON, assets_dir
7+
from ragchecker import RAGChecker, RAGResults
8+
from ragchecker.metrics import all_metrics
9+
10+
QNA_BENCHMARKS = "evaluation/qna_benchmarks.json"
11+
INTERMIDATE_RESULTS = "evaluation/qna_benchmarks_answered.json"
12+
13+
14+
def answer_queries(qa: dict) -> dict:
15+
"""Answer test questions in the RAG benchmark using the Chatbot.
16+
17+
Parameters
18+
----------
19+
qa : dict
20+
A dictionary containing the benchmark questions and answers.
21+
22+
{
23+
"results": [ # A list of QA pairs
24+
{
25+
"query_id": "000",
26+
"query": "This is the question for the first example",
27+
"gt_answer": "This is the ground truth answer for the first example"
28+
},
29+
{
30+
"query_id": "001",
31+
"query": "This is the question for the second example",
32+
"gt_answer": "This is the ground truth answer for the second example"
33+
},
34+
...
35+
]
36+
}
37+
38+
Returns
39+
-------
40+
dict
41+
A dictionary containing the benchmark questions and answers with responses and retrieved
42+
contexts.
43+
44+
{
45+
"results": [ # A list of QA pairs with responses
46+
{
47+
"query_id": "000",
48+
"query": "This is the question for the first example",
49+
"gt_answer": "This is the ground truth answer for the first example",
50+
"response": "This is the response generated by the chatbot",
51+
"retrieved_context": [
52+
{"doc_id": "doc1", "text": "Content from document 1"},
53+
{"doc_id": "doc2", "text": "Content from document 2"}
54+
]
55+
},
56+
...
57+
]
58+
}
59+
60+
"""
61+
if not qa or "results" not in qa:
62+
raise ValueError("Input must contain 'results' key")
63+
64+
chatbot = HybridQueryChatbot(
65+
model_name="gpt-4.1",
66+
model_provider="openai",
67+
paper_path=str(assets_dir / "strainrelief_preprint.pdf"),
68+
references_dir=str(assets_dir / "references"),
69+
)
70+
71+
for qa_pair in qa["results"]:
72+
if "query" not in qa_pair.keys() or "gt_answer" not in qa_pair.keys():
73+
raise ValueError("Each QA pair must contain 'query' and 'gt_answer' keys")
74+
75+
# Stream the response (consuming all chunks)
76+
for _ in chatbot.stream_response(qa_pair["query"]):
77+
pass
78+
79+
# Extract response and context from chat history
80+
last_message = chatbot.chat_history[-1]
81+
qa_pair["response"] = last_message.content
82+
qa_pair["retrieved_context"] = _extract_context(last_message)
83+
84+
return qa
85+
86+
87+
def _extract_context(message) -> list[dict]:
88+
"""Extract and format retrieved context from chat message."""
89+
context_data = message.response_metadata.get("context", [])
90+
return [
91+
{"doc_id": context["Document Title"], "text": context["Content"]}
92+
for context in context_data
93+
if "Document Title" in context and "Content" in context
94+
]
95+
96+
97+
if __name__ == "__main__":
98+
# Load the RAG benchmark Q&As from a JSON file
99+
with open(QNA_BENCHMARKS) as f:
100+
qa = json.load(f)
101+
102+
# Answer each query using the Chatbot
103+
qa_answered = answer_queries(qa)
104+
105+
# Save the updated Q&As with responses and retrieved documents
106+
with open(INTERMIDATE_RESULTS, "w") as f:
107+
json.dump(qa_answered, f, indent=4)
108+
109+
# Initialise RAGResults from the answered Q&As
110+
rag_results = RAGResults.from_dict(qa_answered)
111+
112+
# Set up the evaluator
113+
evaluator = RAGChecker(
114+
extractor_name="openai/gpt-4.1",
115+
checker_name="openai/gpt-4.1",
116+
batch_size_extractor=32,
117+
batch_size_checker=32,
118+
)
119+
120+
# Evaluate results with selected metrics or certain groups
121+
# e.g., retriever_metrics, generator_metrics, all_metrics
122+
evaluator.evaluate(rag_results, all_metrics, save_path=METRICS_JSON)
123+
124+
logger.info(rag_results)
125+
logger.info(f"Evaluation complete. Metrics saved to {METRICS_JSON}.")

evaluation/qna_benchmarks.json

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,14 @@
1+
{
2+
"results": [
3+
{
4+
"query_id": "000",
5+
"query": "This is the question for the first example",
6+
"gt_answer": "This is the ground truth answer for the first example"
7+
},
8+
{
9+
"query_id": "001",
10+
"query": "This is the question for the second example",
11+
"gt_answer": "This is the ground truth answer for the second example"
12+
}
13+
]
14+
}

requirements.in

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,3 +9,4 @@ chromadb
99
sentence-transformers
1010
accelerate
1111
loguru
12+
ragchecker

src/paper_query/chatbots/_chatbots.py

Lines changed: 41 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,9 @@ def __init__(self, model_name: str, model_provider: str):
3131
self.model = setup_model(model_name, model_provider)
3232
self.chain = setup_chain(self.model, prompt=base_prompt)
3333

34-
def stream_response(self, user_input: str, chain_args: dict = {}) -> Generator[str, None, None]:
34+
def stream_response(
35+
self, user_input: str, chain_args: dict = {}, metadata: dict | None = None
36+
) -> Generator[str, None, None]:
3537
"""Process user input and stream AI response."""
3638
# Add user message to history before streaming
3739
logger.debug(f'User input:\n"{user_input}"')
@@ -45,7 +47,11 @@ def stream_response(self, user_input: str, chain_args: dict = {}) -> Generator[s
4547
yield chunk
4648

4749
# After streaming is complete, add the full response to chat history
48-
self.chat_history.append(AIMessage(content=full_response))
50+
self.chat_history.append(
51+
AIMessage(
52+
content=full_response, response_metadata={"context": metadata} if metadata else {}
53+
)
54+
)
4955
logger.debug(f'AI response:\n"{full_response}"')
5056

5157

@@ -121,6 +127,13 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
121127
relevant_references = "\n".join(
122128
[f"From {doc.metadata[RAG_DOC_ID]}:\n{doc.page_content}" for doc in relevant_docs]
123129
)
130+
relevant_metadata = [
131+
{
132+
"Document Title": doc.metadata.get(RAG_DOC_ID, "N/A"),
133+
"Content": doc.page_content,
134+
}
135+
for doc in relevant_docs
136+
]
124137

125138
# Log the context documents
126139
logger.debug(f"Context: {len(relevant_docs)} documents returned.")
@@ -133,7 +146,9 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
133146
)
134147

135148
return super().stream_response(
136-
user_input, {"paper_text": self.paper_text, "relevant_references": relevant_references}
149+
user_input,
150+
{"paper_text": self.paper_text, "relevant_references": relevant_references},
151+
relevant_metadata,
137152
)
138153

139154

@@ -189,6 +204,13 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
189204
relevant_code = "\n".join(
190205
[f"From {doc.metadata[RAG_DOC_ID]}:\n{doc.page_content}" for doc in relevant_docs]
191206
)
207+
relevant_metadata = [
208+
{
209+
"Document Title": doc.metadata.get(RAG_DOC_ID, "N/A"),
210+
"Content": doc.page_content,
211+
}
212+
for doc in relevant_docs
213+
]
192214

193215
# Log the context documents
194216
logger.debug(f"Context: {len(relevant_docs)} documents returned.")
@@ -201,7 +223,9 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
201223
)
202224

203225
return super().stream_response(
204-
user_input, {"paper_text": self.paper_text, "relevant_code": relevant_code}
226+
user_input,
227+
{"paper_text": self.paper_text, "relevant_code": relevant_code},
228+
relevant_metadata,
205229
)
206230

207231

@@ -264,14 +288,22 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
264288
relevant_references = "\n".join(
265289
[f"From {doc.metadata[RAG_DOC_ID]}:\n{doc.page_content}" for doc in relevant_docs]
266290
)
291+
relevant_metadata = [
292+
{
293+
"Document Title": doc.metadata.get(RAG_DOC_ID, "N/A"),
294+
"Content": doc.page_content,
295+
}
296+
for doc in relevant_docs
297+
]
267298

268299
# Log the context documents
269300
logger.debug(f"Context: {len(relevant_docs)} documents returned.")
270-
for i, doc in enumerate(relevant_docs, start=1):
271-
contents = doc.page_content[:200].replace("\n", " ")
301+
for i, doc in enumerate(relevant_metadata, start=1):
302+
contents = doc["Content"][:200].replace("\n", " ")
272303
logger.debug(
273-
f"""Context Document {i}:\nDocument Title: {doc.metadata.get(RAG_DOC_ID, "N/A")}
274-
Page Content: {contents}...
304+
f"""Context Document {i}:
305+
Document Title: {doc["Document Title"]}
306+
Page Content: {contents}
275307
"""
276308
)
277309

@@ -281,4 +313,5 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
281313
"paper_text": self.paper_text,
282314
"relevant_references": relevant_references,
283315
},
316+
relevant_metadata,
284317
)

src/paper_query/constants/__init__.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,21 @@
11
from ._api_keys import GROQ_API_KEY, HUGGINGFACE_API_KEY, OPENAI_API_KEY
2-
from ._paths import PERSIST_DIRECTORY, assets_dir, data_dir, project_dir, src_dir, test_dir
2+
from ._paths import (
3+
METRICS_JSON,
4+
PERSIST_DIRECTORY,
5+
assets_dir,
6+
data_dir,
7+
project_dir,
8+
src_dir,
9+
test_dir,
10+
)
311
from ._strings import RAG_DOC_ID, STREAMLIT_CHEAP_MODEL, STREAMLIT_EXPENSIVE_MODEL
412

513
__all__ = [
614
"OPENAI_API_KEY",
715
"HUGGINGFACE_API_KEY",
816
"GROQ_API_KEY",
917
"PERSIST_DIRECTORY",
18+
"METRICS_JSON",
1019
"project_dir",
1120
"src_dir",
1221
"test_dir",

src/paper_query/constants/_paths.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,3 +8,4 @@
88
assets_dir: Path = project_dir / "assets"
99

1010
PERSIST_DIRECTORY: str = str(project_dir / "vectorstore")
11+
METRICS_JSON: str = str(project_dir / "evaluation" / "rag_evaluation_results.json")

src/paper_query/data/loaders.py

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,13 +10,24 @@
1010
from paper_query.llm import setup_model
1111

1212

13-
def pypdf_loader(file_path: str) -> Document:
13+
def pypdf_loader(file_path: str, interpret_images: bool = False, **image_kwargs) -> Document:
14+
"""Function to load a PDF file, optionally interpreting images."""
15+
if interpret_images and "model" not in image_kwargs:
16+
raise ValueError("When interpret_images is True, 'model' must be provided in image_kwargs.")
17+
18+
if interpret_images:
19+
return _pypdf_loader_w_images(file_path, **image_kwargs)
20+
else:
21+
return _pypdf_loader(file_path)
22+
23+
24+
def _pypdf_loader(file_path: str) -> Document:
1425
"""Function to load text from a PDF file."""
1526
logger.debug("Loading PDF file using PyPDFLoader")
1627
return PyPDFLoader(file_path, mode="single").load()[0]
1728

1829

19-
def pypdf_loader_w_images(
30+
def _pypdf_loader_w_images(
2031
file_path: str, model: str, provider: str, max_tokens: int = 1024
2132
) -> Document:
2233
"""Function to load text from a PDF file with images."""

src/paper_query/llm/models.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,11 @@
11
import os
22

33
from langchain.chat_models import init_chat_model
4+
from langchain_core.language_models.chat_models import BaseChatModel
45
from loguru import logger
56

67

7-
def setup_model(model_name: str, model_provider: str, **kwargs):
8+
def setup_model(model_name: str, model_provider: str, **kwargs) -> BaseChatModel:
89
"""Initialize the chat model."""
910
logger.info(f"Initializing {model_name} model from {model_provider}")
1011
if model_provider == "openai":

src/paper_query/ui/strain_relief_app.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ def strain_relief_chatbot():
3030
"""Chatbot for the StrainRelief paper."""
3131
initialize_session_state()
3232

33-
st.title("The StrainRelief Chatbot")
33+
st.title("StrainReliefChat")
3434
chat_tab, about_tab = st.tabs(["Chat", "About"])
3535

3636
st.sidebar.title("API Configuration")
@@ -46,12 +46,14 @@ def strain_relief_chatbot():
4646
# Display current model
4747
st.sidebar.markdown(f"Using **{st.session_state.model_name}** model.")
4848

49-
st.session_state.chatbot = HybridQueryChatbot(
50-
model_name=st.session_state.model_name.lower(),
51-
model_provider="openai",
52-
paper_path=str(assets_dir / "strainrelief_preprint.pdf"),
53-
references_dir=str(assets_dir / "references"),
54-
)
49+
# Only instantiate chatbot once and store in session state
50+
if st.session_state.chatbot is None:
51+
st.session_state.chatbot = HybridQueryChatbot(
52+
model_name=st.session_state.model_name.lower(),
53+
model_provider="openai",
54+
paper_path=str(assets_dir / "strainrelief_preprint.pdf"),
55+
references_dir=str(assets_dir / "references"),
56+
)
5557

5658
with chat_tab:
5759
if "messages" not in st.session_state:

0 commit comments

Comments
 (0)