chATLAS_Chains 0.1.5__tar.gz → 0.1.6__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/PKG-INFO +7 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/README.md +6 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/advanced.py +32 -35
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rrf.py +29 -26
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/groq.py +2 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/model_selection.py +9 -5
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/PKG-INFO +7 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/SOURCES.txt +1 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/pyproject.toml +1 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_chains.py +31 -19
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_llm.py +26 -4
- chatlas_chains-0.1.6/tests/test_rrf.py +126 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/LICENSE +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic_graph.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rerank.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/log.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/starters.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/query/query_rewriting.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/search/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/search/basic.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/doc_utils.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/vectorstore.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/requires.txt +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/top_level.txt +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/setup.cfg +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/conftest.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_groq.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_search.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: chATLAS_Chains
|
|
3
|
-
Version: 0.1.
|
|
3
|
+
Version: 0.1.6
|
|
4
4
|
Summary: A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project.
|
|
5
5
|
Author-email: Joe Egan <joseph.caimin.egan@cern.ch>
|
|
6
6
|
License: Apache-2.0
|
|
@@ -91,6 +91,12 @@ CREATE EXTENSION IF NOT EXISTS vector;
|
|
|
91
91
|
```
|
|
92
92
|
## CHANGELOG
|
|
93
93
|
|
|
94
|
+
#### 0.1.6
|
|
95
|
+
|
|
96
|
+
Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
|
|
97
|
+
|
|
98
|
+
Add `fallback_models` optional argument to `advanced_rag`
|
|
99
|
+
|
|
94
100
|
#### 0.1.5
|
|
95
101
|
|
|
96
102
|
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
@@ -65,6 +65,12 @@ CREATE EXTENSION IF NOT EXISTS vector;
|
|
|
65
65
|
```
|
|
66
66
|
## CHANGELOG
|
|
67
67
|
|
|
68
|
+
#### 0.1.6
|
|
69
|
+
|
|
70
|
+
Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
|
|
71
|
+
|
|
72
|
+
Add `fallback_models` optional argument to `advanced_rag`
|
|
73
|
+
|
|
68
74
|
#### 0.1.5
|
|
69
75
|
|
|
70
76
|
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
@@ -11,6 +11,7 @@ upweighting results that appear in both
|
|
|
11
11
|
"""
|
|
12
12
|
|
|
13
13
|
import argparse
|
|
14
|
+
import logging
|
|
14
15
|
import os
|
|
15
16
|
import sys
|
|
16
17
|
from typing import TypedDict
|
|
@@ -30,6 +31,8 @@ from chATLAS_Chains.search.basic import search_runnable
|
|
|
30
31
|
from chATLAS_Chains.utils.doc_utils import combine_documents
|
|
31
32
|
from chATLAS_Embed.Base import VectorStore
|
|
32
33
|
|
|
34
|
+
logger = logging.getLogger(__name__)
|
|
35
|
+
|
|
33
36
|
|
|
34
37
|
# Define TypedDict for the simplified state
|
|
35
38
|
class HybridGraphState(TypedDict, total=False):
|
|
@@ -57,6 +60,7 @@ def advanced_rag(
|
|
|
57
60
|
rrf_constant: float = 60.0,
|
|
58
61
|
rrf_weights: dict[str, float] | None = None,
|
|
59
62
|
retry_config: RetryConfig | None = None,
|
|
63
|
+
fallback_models: list[str] | None = None,
|
|
60
64
|
) -> CompiledStateGraph:
|
|
61
65
|
"""
|
|
62
66
|
Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
|
|
@@ -77,6 +81,7 @@ def advanced_rag(
|
|
|
77
81
|
:param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
|
|
78
82
|
:param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
|
|
79
83
|
:param retry_config: Optional RetryConfig for Groq API calls
|
|
84
|
+
:param fallback_models: Optional list of model names to fall back to if the Groq API fails.
|
|
80
85
|
|
|
81
86
|
:return: A compiled LangGraph with the chosen features
|
|
82
87
|
|
|
@@ -91,13 +96,15 @@ def advanced_rag(
|
|
|
91
96
|
searcher = search_runnable(vectorstore)
|
|
92
97
|
|
|
93
98
|
# Initialise models
|
|
94
|
-
model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
|
|
99
|
+
model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config, fallback_models)
|
|
95
100
|
|
|
96
101
|
if enable_query_rewriting:
|
|
97
102
|
query_rewriting_model_instance = get_chat_model(
|
|
98
103
|
model_name=query_rewriting_model,
|
|
99
104
|
temperature=query_rewriting_temperature,
|
|
100
105
|
use_preview_models=use_preview_models,
|
|
106
|
+
retry_config=retry_config,
|
|
107
|
+
fallback_models=fallback_models,
|
|
101
108
|
)
|
|
102
109
|
|
|
103
110
|
# --------------- Define functions for the graph nodes ---------------
|
|
@@ -107,7 +114,7 @@ def advanced_rag(
|
|
|
107
114
|
question = state["question"]
|
|
108
115
|
|
|
109
116
|
if not enable_query_rewriting:
|
|
110
|
-
|
|
117
|
+
logger.warning("Query rewriting disabled, but query_rewriting node was called.")
|
|
111
118
|
# return original state
|
|
112
119
|
return {**state}
|
|
113
120
|
|
|
@@ -128,7 +135,7 @@ def advanced_rag(
|
|
|
128
135
|
|
|
129
136
|
search_kwargs = state.get("search_kwargs", {})
|
|
130
137
|
if not search_kwargs:
|
|
131
|
-
|
|
138
|
+
logger.warning("No search_kwargs provided, using defaults.")
|
|
132
139
|
results = searcher.invoke(question)
|
|
133
140
|
else:
|
|
134
141
|
results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
|
|
@@ -136,7 +143,7 @@ def advanced_rag(
|
|
|
136
143
|
if "docs" not in results:
|
|
137
144
|
raise Exception("Search results missing 'docs' field")
|
|
138
145
|
|
|
139
|
-
|
|
146
|
+
logger.debug(f"Retrieved {len(results['docs'])} documents")
|
|
140
147
|
|
|
141
148
|
return {
|
|
142
149
|
**state,
|
|
@@ -150,7 +157,7 @@ def advanced_rag(
|
|
|
150
157
|
"""
|
|
151
158
|
|
|
152
159
|
if not enable_reranking:
|
|
153
|
-
|
|
160
|
+
logger.warning("Reranking disabled, but rerank node was called.")
|
|
154
161
|
# return original state
|
|
155
162
|
return {**state}
|
|
156
163
|
|
|
@@ -166,7 +173,7 @@ def advanced_rag(
|
|
|
166
173
|
)
|
|
167
174
|
|
|
168
175
|
except Exception as e:
|
|
169
|
-
|
|
176
|
+
logger.warning(f"Reranking failed with exception {e}. Returning original documents.")
|
|
170
177
|
reranked_docs = docs
|
|
171
178
|
|
|
172
179
|
return {
|
|
@@ -177,13 +184,13 @@ def advanced_rag(
|
|
|
177
184
|
def rrf(state: HybridGraphState) -> HybridGraphState:
|
|
178
185
|
"""Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
|
|
179
186
|
if not enable_rrf:
|
|
180
|
-
|
|
187
|
+
logger.warning("RRF disabled, but rrf node was called.")
|
|
181
188
|
# return original state
|
|
182
189
|
return {**state}
|
|
183
190
|
|
|
184
191
|
docs = state.get("docs", [])
|
|
185
192
|
if not docs:
|
|
186
|
-
|
|
193
|
+
logger.warning("No documents retrieved for RRF.")
|
|
187
194
|
return {**state}
|
|
188
195
|
|
|
189
196
|
try:
|
|
@@ -191,7 +198,7 @@ def advanced_rag(
|
|
|
191
198
|
results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
|
|
192
199
|
)
|
|
193
200
|
except Exception as e:
|
|
194
|
-
|
|
201
|
+
logger.warning(f"RRF failed with exception {e}. Returning original documents.")
|
|
195
202
|
rrf_docs = docs
|
|
196
203
|
|
|
197
204
|
return {
|
|
@@ -201,7 +208,6 @@ def advanced_rag(
|
|
|
201
208
|
|
|
202
209
|
def generate_answer(state: HybridGraphState) -> HybridGraphState:
|
|
203
210
|
"""Generate a response from the LLM"""
|
|
204
|
-
|
|
205
211
|
# Format the prompt using LangChain template
|
|
206
212
|
prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
|
|
207
213
|
final_prompt = prompt_template.format_messages(**prompt_input)
|
|
@@ -266,6 +272,7 @@ if __name__ == "__main__":
|
|
|
266
272
|
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
267
273
|
|
|
268
274
|
twiki = get_vectorstore("twiki_prod")
|
|
275
|
+
mkdocs = get_vectorstore("mkdocs_prod_v1")
|
|
269
276
|
|
|
270
277
|
retry_config = RetryConfig(
|
|
271
278
|
max_retries=1,
|
|
@@ -274,34 +281,24 @@ if __name__ == "__main__":
|
|
|
274
281
|
|
|
275
282
|
# Create the hybrid graph
|
|
276
283
|
graph = advanced_rag(
|
|
277
|
-
vectorstore=[twiki],
|
|
284
|
+
vectorstore=[twiki, mkdocs],
|
|
278
285
|
model_name=GROQ_PRODUCTION_MODELS[0],
|
|
279
286
|
enable_query_rewriting=True,
|
|
280
287
|
enable_rrf=True,
|
|
281
|
-
enable_reranking=
|
|
288
|
+
enable_reranking=False,
|
|
282
289
|
)
|
|
283
290
|
|
|
284
|
-
|
|
285
|
-
|
|
286
|
-
|
|
287
|
-
{
|
|
288
|
-
"
|
|
289
|
-
"
|
|
290
|
-
|
|
291
|
-
|
|
292
|
-
|
|
293
|
-
|
|
294
|
-
|
|
295
|
-
}
|
|
296
|
-
)
|
|
297
|
-
|
|
298
|
-
print(f"Number of docs is : {len(ans['docs'])}")
|
|
299
|
-
print(f"Answer: {ans['answer']}")
|
|
300
|
-
|
|
301
|
-
except Exception as e:
|
|
302
|
-
print(f"❌ Graph execution failed with error: {e}")
|
|
303
|
-
print(f"Error type: {type(e).__name__}")
|
|
304
|
-
import traceback
|
|
291
|
+
ans = graph.invoke(
|
|
292
|
+
{
|
|
293
|
+
"question": "What is the crack veto in electron reconstruction?",
|
|
294
|
+
"search_kwargs": {
|
|
295
|
+
"k_text": 3,
|
|
296
|
+
"k": 15,
|
|
297
|
+
"date_filter": "01-01-2010",
|
|
298
|
+
# "type": ["twiki"],
|
|
299
|
+
},
|
|
300
|
+
}
|
|
301
|
+
)
|
|
305
302
|
|
|
306
|
-
|
|
307
|
-
|
|
303
|
+
print(f"Number of docs is : {len(ans['docs'])}")
|
|
304
|
+
print(f"Answer: {ans['answer']}")
|
|
@@ -5,7 +5,9 @@ Functions for performing Reciprocal Rank Fusion (RRF) on documents retrieved fro
|
|
|
5
5
|
from langchain_core.documents import Document
|
|
6
6
|
|
|
7
7
|
|
|
8
|
-
def reciprocal_rank_fusion(
|
|
8
|
+
def reciprocal_rank_fusion(
|
|
9
|
+
results: dict[str, list[Document]], k: float = 60, weights: dict[str, float] | None = None
|
|
10
|
+
) -> list:
|
|
9
11
|
"""
|
|
10
12
|
Fuse search results using weighted Reciprocal Rank Fusion (RRF) algorithm.
|
|
11
13
|
|
|
@@ -39,24 +41,21 @@ def reciprocal_rank_fusion(results: dict, k: float = 60, weights: dict[str, floa
|
|
|
39
41
|
if key not in weights:
|
|
40
42
|
weights[key] = 1.0
|
|
41
43
|
|
|
42
|
-
if len(results) != len(weights):
|
|
43
|
-
raise ValueError("Number of results and weights are different")
|
|
44
|
-
|
|
45
44
|
# Do this for each retriever in the dict
|
|
46
45
|
for retriever, docs in results.items():
|
|
47
46
|
for rank, doc in enumerate(docs, 1):
|
|
48
|
-
|
|
47
|
+
name = doc.metadata.get("name")
|
|
49
48
|
|
|
50
49
|
# initialise score to zero
|
|
51
|
-
if
|
|
52
|
-
rrf_scores[
|
|
53
|
-
document_store[
|
|
50
|
+
if name not in rrf_scores:
|
|
51
|
+
rrf_scores[name] = 0
|
|
52
|
+
document_store[name] = doc
|
|
54
53
|
|
|
55
54
|
# weighted score for this retriever
|
|
56
|
-
rrf_scores[
|
|
55
|
+
rrf_scores[name] += weights[retriever] * (1 / (k + rank))
|
|
57
56
|
|
|
58
57
|
# track which retrievers had this document
|
|
59
|
-
metadata = document_store[
|
|
58
|
+
metadata = document_store[name].metadata
|
|
60
59
|
if "retrievers" not in metadata:
|
|
61
60
|
metadata["retrievers"] = []
|
|
62
61
|
if retriever not in metadata["retrievers"]:
|
|
@@ -64,9 +63,10 @@ def reciprocal_rank_fusion(results: dict, k: float = 60, weights: dict[str, floa
|
|
|
64
63
|
|
|
65
64
|
# Sort and return results
|
|
66
65
|
sorted_results = []
|
|
67
|
-
|
|
68
|
-
|
|
69
|
-
doc
|
|
66
|
+
|
|
67
|
+
for name in sorted(rrf_scores.keys(), key=lambda x: rrf_scores[x], reverse=True):
|
|
68
|
+
doc = document_store[name]
|
|
69
|
+
doc.metadata["rrf_score"] = rrf_scores[name]
|
|
70
70
|
|
|
71
71
|
sorted_results.append(doc)
|
|
72
72
|
|
|
@@ -90,7 +90,6 @@ def split_docs_by_retriever(docs: list[Document]):
|
|
|
90
90
|
retriever_types.append(search_type)
|
|
91
91
|
|
|
92
92
|
split_docs = {retriever_type: [] for retriever_type in retriever_types}
|
|
93
|
-
|
|
94
93
|
for doc in docs:
|
|
95
94
|
search_type = doc.metadata["search_type"]
|
|
96
95
|
if search_type not in split_docs:
|
|
@@ -102,36 +101,40 @@ def split_docs_by_retriever(docs: list[Document]):
|
|
|
102
101
|
|
|
103
102
|
if __name__ == "__main__":
|
|
104
103
|
# Example usage
|
|
105
|
-
from langchain_core.documents import Document
|
|
106
104
|
|
|
107
105
|
# Create some example documents
|
|
108
106
|
doc1 = Document(
|
|
109
107
|
page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
|
|
110
|
-
|
|
111
|
-
|
|
108
|
+
metadata={"search_type": "vector", "name": "higgs"},
|
|
109
|
+
source="twiki",
|
|
110
|
+
url="",
|
|
112
111
|
)
|
|
113
112
|
doc2 = Document(
|
|
114
113
|
page_content="The top quark is the heaviest known elementary particle",
|
|
115
|
-
|
|
116
|
-
|
|
114
|
+
metadata={"search_type": "vector", "name": "top"},
|
|
115
|
+
source="twiki",
|
|
116
|
+
url="",
|
|
117
117
|
)
|
|
118
118
|
|
|
119
119
|
vector_results = [doc1, doc2]
|
|
120
120
|
|
|
121
121
|
doc3 = Document(
|
|
122
122
|
page_content="The gluon is the force carrier of Quantum Chromodynamics (QCD)",
|
|
123
|
-
|
|
124
|
-
|
|
123
|
+
metadata={"search_type": "text", "name": "gluon"},
|
|
124
|
+
source="twiki",
|
|
125
|
+
url="",
|
|
125
126
|
)
|
|
126
127
|
doc4 = Document(
|
|
127
128
|
page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
|
|
128
|
-
|
|
129
|
-
|
|
129
|
+
metadata={"search_type": "text", "name": "higgs"},
|
|
130
|
+
source="twiki",
|
|
131
|
+
url="",
|
|
130
132
|
)
|
|
131
133
|
doc5 = Document(
|
|
132
134
|
page_content="The Large Hadron Collider (LHC) is the world's largest and most powerful particle accelerator",
|
|
133
|
-
|
|
134
|
-
|
|
135
|
+
metadata={"search_type": "text", "name": "lhc"},
|
|
136
|
+
source="twiki",
|
|
137
|
+
url="",
|
|
135
138
|
)
|
|
136
139
|
|
|
137
140
|
text_results = [doc3, doc4, doc5]
|
|
@@ -143,5 +146,5 @@ if __name__ == "__main__":
|
|
|
143
146
|
print("RRF Results:")
|
|
144
147
|
for doc in rrf_results:
|
|
145
148
|
print(
|
|
146
|
-
f"Document
|
|
149
|
+
f"Document name: {doc.metadata['name']}, RRF Score: {doc.metadata['rrf_score']}, Retrievers: {doc.metadata['retrievers']}"
|
|
147
150
|
)
|
|
@@ -139,6 +139,7 @@ class AccGPTChatGroq(BaseChatModel, BaseModel):
|
|
|
139
139
|
def __init__(self, **data):
|
|
140
140
|
super().__init__(**data)
|
|
141
141
|
if self.retry_config is None:
|
|
142
|
+
logger.debug("No retry_config provided, using default")
|
|
142
143
|
self.retry_config = RetryConfig()
|
|
143
144
|
|
|
144
145
|
if self.fallback_models is None:
|
|
@@ -565,7 +566,7 @@ if __name__ == "__main__":
|
|
|
565
566
|
api_key = api_key.strip()
|
|
566
567
|
|
|
567
568
|
base_url = os.getenv("CHATLAS_GROQ_BASE_URL")
|
|
568
|
-
base_url = "http://localhost:3000"
|
|
569
|
+
# base_url = "http://localhost:3000"
|
|
569
570
|
if not base_url or not base_url.strip():
|
|
570
571
|
raise ValueError("CHATLAS_GROQ_BASE_URL not set in environment")
|
|
571
572
|
base_url = base_url.strip()
|
|
@@ -84,6 +84,7 @@ def get_chat_model(
|
|
|
84
84
|
temperature: float | None = None,
|
|
85
85
|
use_preview_models: bool = False,
|
|
86
86
|
retry_config: RetryConfig | None = None,
|
|
87
|
+
fallback_models: list[str] | None = None,
|
|
87
88
|
):
|
|
88
89
|
"""
|
|
89
90
|
Initialize chat model with the provided model name (if supported)
|
|
@@ -100,6 +101,12 @@ def get_chat_model(
|
|
|
100
101
|
:param use_preview_models: If True, allows the use of preview models from Groq. Defaults to False.
|
|
101
102
|
:type use_preview_models: bool
|
|
102
103
|
|
|
104
|
+
:param retry_config: Configuration for retrying requests to Groq models. If None, the default retry configuration is used.
|
|
105
|
+
:type retry_config: RetryConfig | None
|
|
106
|
+
|
|
107
|
+
:param fallback_models: A list of model names to fall back to if the Groq API fails.
|
|
108
|
+
:type fallback_models: list[str] | None
|
|
109
|
+
|
|
103
110
|
:raises ValueError: If the environment variable `CHATLAS_OPENAI_KEY` is not set when using OpenAI models, or if `CHATLAS_GROQ_KEY` and `CHATLAS_GROQ_BASE_URL` are not set when using Groq models.
|
|
104
111
|
|
|
105
112
|
:return: An instance of the specified chat model.
|
|
@@ -143,17 +150,14 @@ def get_chat_model(
|
|
|
143
150
|
f"max_tokens ({max_tokens}) exceeds the model's maximum ({get_max_completion_tokens(model_name)})."
|
|
144
151
|
)
|
|
145
152
|
|
|
146
|
-
# default retry config
|
|
147
|
-
if retry_config is None:
|
|
148
|
-
retry_config = RetryConfig()
|
|
149
|
-
|
|
150
153
|
llm = AccGPTChatGroq(
|
|
151
154
|
model_name=model_name,
|
|
152
155
|
api_key=api_key,
|
|
153
156
|
base_url=base_url,
|
|
154
157
|
max_tokens=max_tokens,
|
|
155
158
|
temperature=temperature,
|
|
156
|
-
retry_config=retry_config,
|
|
159
|
+
retry_config=retry_config, # if None, AccGPTChatGroq will use the default
|
|
160
|
+
fallback_models=fallback_models,
|
|
157
161
|
)
|
|
158
162
|
|
|
159
163
|
return llm
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: chATLAS_Chains
|
|
3
|
-
Version: 0.1.
|
|
3
|
+
Version: 0.1.6
|
|
4
4
|
Summary: A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project.
|
|
5
5
|
Author-email: Joe Egan <joseph.caimin.egan@cern.ch>
|
|
6
6
|
License: Apache-2.0
|
|
@@ -91,6 +91,12 @@ CREATE EXTENSION IF NOT EXISTS vector;
|
|
|
91
91
|
```
|
|
92
92
|
## CHANGELOG
|
|
93
93
|
|
|
94
|
+
#### 0.1.6
|
|
95
|
+
|
|
96
|
+
Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
|
|
97
|
+
|
|
98
|
+
Add `fallback_models` optional argument to `advanced_rag`
|
|
99
|
+
|
|
94
100
|
#### 0.1.5
|
|
95
101
|
|
|
96
102
|
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "chATLAS_Chains"
|
|
7
|
-
version = "0.1.
|
|
7
|
+
version = "0.1.6"
|
|
8
8
|
description = "A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project."
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
requires-python = ">=3.11"
|
|
@@ -22,7 +22,7 @@ from chATLAS_Chains.vectorstore import get_vectorstore
|
|
|
22
22
|
@pytest.fixture
|
|
23
23
|
def retry_config():
|
|
24
24
|
"""Retry config to use in tests"""
|
|
25
|
-
return RetryConfig(max_retries=
|
|
25
|
+
return RetryConfig(max_retries=5, max_delay=120.0)
|
|
26
26
|
|
|
27
27
|
|
|
28
28
|
def test_basic_retrieval_chain_returns_runnablesequence(twiki_vectorstore):
|
|
@@ -55,10 +55,11 @@ def test_basic_retrieval_chain_returns_docs_and_answer(twiki_vectorstore, retry_
|
|
|
55
55
|
)
|
|
56
56
|
output = chain.invoke("What is the Higgs boson?")
|
|
57
57
|
|
|
58
|
-
assert "docs" in output
|
|
59
|
-
assert "answer" in output
|
|
58
|
+
assert "docs" in output, "'docs' key not in output"
|
|
59
|
+
assert "answer" in output, "'answer' key not in output"
|
|
60
60
|
|
|
61
|
-
assert isinstance(output["docs"], list)
|
|
61
|
+
assert isinstance(output["docs"], list), "'docs' is not a list"
|
|
62
|
+
assert all([isinstance(obj, Document) for obj in output["docs"]]), "'docs' should be a list of Documents"
|
|
62
63
|
|
|
63
64
|
|
|
64
65
|
def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores, retry_config):
|
|
@@ -70,10 +71,11 @@ def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores, retry_c
|
|
|
70
71
|
)
|
|
71
72
|
output = chain.invoke("What is the Higgs boson?")
|
|
72
73
|
|
|
73
|
-
assert "docs" in output
|
|
74
|
-
assert "
|
|
74
|
+
assert "docs" in output, "'docs' key not in output"
|
|
75
|
+
assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3, three vectorstores"
|
|
76
|
+
assert "answer" in output, "'answer' key not in output"
|
|
75
77
|
|
|
76
|
-
assert isinstance(output["docs"], list)
|
|
78
|
+
assert isinstance(output["docs"], list), "'docs' is not a list"
|
|
77
79
|
|
|
78
80
|
sources = set([doc.metadata.get("source") for doc in output["docs"]])
|
|
79
81
|
expected_sources = {"twiki", "MkDocs", "CDS"}
|
|
@@ -101,10 +103,13 @@ def test_advanced_rag_chain(twiki_vectorstore, retry_config):
|
|
|
101
103
|
{"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
|
|
102
104
|
)
|
|
103
105
|
|
|
104
|
-
assert "docs" in output
|
|
105
|
-
assert "
|
|
106
|
-
assert
|
|
107
|
-
assert isinstance(output["docs"]
|
|
106
|
+
assert "docs" in output, "'docs' key not in output"
|
|
107
|
+
assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3"
|
|
108
|
+
assert "answer" in output, "'answer' key not in output"
|
|
109
|
+
assert isinstance(output["docs"], list), "'docs' is not a list"
|
|
110
|
+
assert isinstance(output["docs"][0], Document), (
|
|
111
|
+
f"'docs' should contain instances of {Document}, but instead it has {type(output['docs'][0])}"
|
|
112
|
+
)
|
|
108
113
|
|
|
109
114
|
# test with the additional options turned on (besides rerank)
|
|
110
115
|
chain = advanced_rag(
|
|
@@ -119,10 +124,14 @@ def test_advanced_rag_chain(twiki_vectorstore, retry_config):
|
|
|
119
124
|
{"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
|
|
120
125
|
)
|
|
121
126
|
|
|
122
|
-
assert "docs" in output
|
|
123
|
-
assert "
|
|
124
|
-
assert
|
|
125
|
-
assert isinstance(output["docs"]
|
|
127
|
+
assert "docs" in output, "'docs' key not in output"
|
|
128
|
+
assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3"
|
|
129
|
+
assert "answer" in output, "'answer' key not in output"
|
|
130
|
+
assert isinstance(output["docs"], list), "'docs' is not a list"
|
|
131
|
+
assert all(isinstance(obj, Document) for obj in output["docs"]), (
|
|
132
|
+
f"'docs' should only contain instances of {Document}"
|
|
133
|
+
)
|
|
134
|
+
assert all("rrf_score" in doc.metadata for doc in output["docs"]), "'rrf_score' not in document metadata"
|
|
126
135
|
|
|
127
136
|
|
|
128
137
|
def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores, retry_config):
|
|
@@ -143,10 +152,13 @@ def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores, retry_conf
|
|
|
143
152
|
{"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
|
|
144
153
|
)
|
|
145
154
|
|
|
146
|
-
assert "docs" in output
|
|
147
|
-
assert "
|
|
148
|
-
assert
|
|
149
|
-
assert isinstance(output["docs"]
|
|
155
|
+
assert "docs" in output, "'docs' key not in output"
|
|
156
|
+
assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3, three vectorstores"
|
|
157
|
+
assert "answer" in output, "'answer' key not in output"
|
|
158
|
+
assert isinstance(output["docs"], list), "'docs' is not a list"
|
|
159
|
+
assert isinstance(output["docs"][0], Document), (
|
|
160
|
+
f"'docs' should contain instances of {Document}, but instead it has {type(output['docs'][0])}"
|
|
161
|
+
)
|
|
150
162
|
|
|
151
163
|
sources = set([doc.metadata.get("source") for doc in output["docs"]])
|
|
152
164
|
expected_sources = {"twiki", "MkDocs", "CDS"}
|
|
@@ -60,7 +60,9 @@ class TestModelSelection:
|
|
|
60
60
|
"""Test that Groq models are correctly initialized."""
|
|
61
61
|
mock_class, mock_instance = mock_groq
|
|
62
62
|
|
|
63
|
-
|
|
63
|
+
retry_config = RetryConfig(max_retries=5, base_delay=5.0, max_delay=300.0)
|
|
64
|
+
|
|
65
|
+
model = get_chat_model(model_name, use_preview_models=True, retry_config=retry_config)
|
|
64
66
|
|
|
65
67
|
mock_class.assert_called_once_with(
|
|
66
68
|
model_name=model_name,
|
|
@@ -68,7 +70,8 @@ class TestModelSelection:
|
|
|
68
70
|
base_url="http://fake-groq-url.com",
|
|
69
71
|
max_tokens=get_max_completion_tokens(model_name),
|
|
70
72
|
temperature=None,
|
|
71
|
-
retry_config=
|
|
73
|
+
retry_config=retry_config,
|
|
74
|
+
fallback_models=None,
|
|
72
75
|
)
|
|
73
76
|
assert model is mock_instance
|
|
74
77
|
|
|
@@ -97,7 +100,8 @@ class TestModelSelection:
|
|
|
97
100
|
base_url="http://fake-groq-url.com",
|
|
98
101
|
max_tokens=get_max_completion_tokens(preview_model),
|
|
99
102
|
temperature=None,
|
|
100
|
-
retry_config=
|
|
103
|
+
retry_config=None,
|
|
104
|
+
fallback_models=None,
|
|
101
105
|
)
|
|
102
106
|
assert model is mock_instance
|
|
103
107
|
|
|
@@ -197,9 +201,27 @@ class TestModelSelection:
|
|
|
197
201
|
base_url="http://whitespace-url.com",
|
|
198
202
|
max_tokens=get_max_completion_tokens(GROQ_PRODUCTION_MODELS[0]),
|
|
199
203
|
temperature=None,
|
|
200
|
-
retry_config=
|
|
204
|
+
retry_config=None,
|
|
205
|
+
fallback_models=None,
|
|
201
206
|
)
|
|
202
207
|
|
|
203
208
|
def test_supported_models_lists(self):
|
|
204
209
|
"""Test that SUPPORTED_CHAT_MODELS is correctly composed of OPENAI_MODELS and GROQ_MODELS."""
|
|
205
210
|
assert set(SUPPORTED_CHAT_MODELS) == set(OPENAI_MODELS + GROQ_MODELS)
|
|
211
|
+
|
|
212
|
+
def test_retry_config_set_correctly(self):
|
|
213
|
+
"""Test that the retry_config is correctly passed to the Groq model."""
|
|
214
|
+
retry_config = RetryConfig(max_retries=3, base_delay=2.0, max_delay=100.0)
|
|
215
|
+
|
|
216
|
+
model = get_chat_model(GROQ_PRODUCTION_MODELS[0], use_preview_models=True, retry_config=retry_config)
|
|
217
|
+
|
|
218
|
+
assert isinstance(model, AccGPTChatGroq), "Returned model is not an instance of AccGPTChatGroq"
|
|
219
|
+
assert model.retry_config.max_retries == 3, (
|
|
220
|
+
f"max_retries not set correctly on model, it set to {model.retry_config.max_retries}"
|
|
221
|
+
)
|
|
222
|
+
assert model.retry_config.base_delay == 2.0, (
|
|
223
|
+
f"base_delay not set correctly on model, it set to {model.retry_config.base_delay}"
|
|
224
|
+
)
|
|
225
|
+
assert model.retry_config.max_delay == 100.0, (
|
|
226
|
+
f"max_delay not set correctly on model, it set to {model.retry_config.max_delay}"
|
|
227
|
+
)
|
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
import pytest
|
|
2
|
+
from langchain_core.documents import Document
|
|
3
|
+
|
|
4
|
+
from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
@pytest.fixture
|
|
8
|
+
def vector_results():
|
|
9
|
+
doc1 = Document(
|
|
10
|
+
page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
|
|
11
|
+
metadata={"search_type": "vector", "name": "higgs"},
|
|
12
|
+
source="twiki",
|
|
13
|
+
url="",
|
|
14
|
+
)
|
|
15
|
+
doc2 = Document(
|
|
16
|
+
page_content="The top quark is the heaviest known elementary particle",
|
|
17
|
+
# name="top",
|
|
18
|
+
metadata={"search_type": "vector", "name": "top"},
|
|
19
|
+
source="twiki",
|
|
20
|
+
url="",
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
return [doc1, doc2]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@pytest.fixture
|
|
27
|
+
def text_results():
|
|
28
|
+
doc3 = Document(
|
|
29
|
+
page_content="The gluon is the force carrier of Quantum Chromodynamics (QCD)",
|
|
30
|
+
# name="gluon",
|
|
31
|
+
metadata={"search_type": "text", "name": "gluon"},
|
|
32
|
+
source="twiki",
|
|
33
|
+
url="",
|
|
34
|
+
)
|
|
35
|
+
doc4 = Document(
|
|
36
|
+
page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
|
|
37
|
+
# name="higgs",
|
|
38
|
+
metadata={"search_type": "text", "name": "higgs"},
|
|
39
|
+
source="twiki",
|
|
40
|
+
url="",
|
|
41
|
+
)
|
|
42
|
+
doc5 = Document(
|
|
43
|
+
page_content="The Large Hadron Collider (LHC) is the world's largest and most powerful particle accelerator",
|
|
44
|
+
# name="lhc",
|
|
45
|
+
metadata={"search_type": "text", "name": "lhc"},
|
|
46
|
+
source="twiki",
|
|
47
|
+
url="",
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
return [doc3, doc4, doc5]
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
@pytest.fixture
|
|
54
|
+
def docs_dict(vector_results, text_results):
|
|
55
|
+
return {
|
|
56
|
+
"vector": vector_results,
|
|
57
|
+
"text": text_results,
|
|
58
|
+
}
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def test_split_docs_by_retriever(vector_results, text_results):
|
|
62
|
+
# pass it the combined list, it should re-split
|
|
63
|
+
combined = vector_results + text_results
|
|
64
|
+
|
|
65
|
+
split = split_docs_by_retriever(combined)
|
|
66
|
+
|
|
67
|
+
assert isinstance(split, dict), "return type should be dict"
|
|
68
|
+
assert "text" in split, "missing retriever type in returned dict"
|
|
69
|
+
assert "vector" in split, "missing retriever type in returned dict"
|
|
70
|
+
assert len(split["text"]) == len(text_results), "text results length mismatch"
|
|
71
|
+
assert len(split["vector"]) == len(vector_results), "vector results length mismatch"
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def test_split_docs_by_retriever_missing_search_type(vector_results):
|
|
75
|
+
# remove search_type from one doc
|
|
76
|
+
vector_results[0].metadata.pop("search_type")
|
|
77
|
+
|
|
78
|
+
with pytest.raises(ValueError, match="search_type metadata field missing from Document.metadata"):
|
|
79
|
+
split_docs_by_retriever(vector_results)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def test_reciprocal_rank_fusion(docs_dict):
|
|
83
|
+
fused = reciprocal_rank_fusion(docs_dict)
|
|
84
|
+
|
|
85
|
+
assert isinstance(fused, list), "return type should be list"
|
|
86
|
+
assert len(fused) <= 5, "fused results length exceeds top_k"
|
|
87
|
+
assert all(isinstance(doc, Document) for doc in fused), f"all items in fused results should be {Document} instances"
|
|
88
|
+
|
|
89
|
+
assert all("rrf_score" in doc.metadata for doc in fused), "rrf_score missing from Document.metadata"
|
|
90
|
+
|
|
91
|
+
# Check that the documents are ordered by score (highest first)
|
|
92
|
+
scores = [doc.metadata["rrf_score"] for doc in fused]
|
|
93
|
+
assert scores == sorted(scores, reverse=True), "documents are not ordered by rrf_score"
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def test_reciprocal_rank_fusion_with_docs_splitter(vector_results, text_results):
|
|
97
|
+
"""
|
|
98
|
+
Same as above, but check it still works using the actual splitter function
|
|
99
|
+
"""
|
|
100
|
+
fused = reciprocal_rank_fusion(split_docs_by_retriever(vector_results + text_results))
|
|
101
|
+
|
|
102
|
+
assert isinstance(fused, list), "return type should be list"
|
|
103
|
+
assert len(fused) <= 5, "fused results length exceeds top_k"
|
|
104
|
+
assert all(isinstance(doc, Document) for doc in fused), f"all items in fused results should be {Document} instances"
|
|
105
|
+
|
|
106
|
+
assert all("rrf_score" in doc.metadata for doc in fused), "rrf_score missing from Document.metadata"
|
|
107
|
+
|
|
108
|
+
# Check that the documents are ordered by score (highest first)
|
|
109
|
+
scores = [doc.metadata["rrf_score"] for doc in fused]
|
|
110
|
+
assert scores == sorted(scores, reverse=True), "documents are not ordered by rrf_score"
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def test_reciprocal_rank_fusion_raise_for_negative_k(docs_dict):
|
|
114
|
+
with pytest.raises(ValueError, match="k must be positive"):
|
|
115
|
+
reciprocal_rank_fusion(docs_dict, k=-7)
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def test_reciprocal_rank_fusion_raise_for_mismatched_weights(docs_dict):
|
|
119
|
+
weights = {
|
|
120
|
+
"vector": 1.0,
|
|
121
|
+
"text": 1.0,
|
|
122
|
+
"extra": 1.0,
|
|
123
|
+
}
|
|
124
|
+
|
|
125
|
+
with pytest.raises(ValueError, match="Weight 'extra' not found in provided results dictionary."):
|
|
126
|
+
reciprocal_rank_fusion(docs_dict, weights=weights)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/enhanced_agentic_graph.py
RENAMED
|
File without changes
|
{chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/websearch_retrieval_chain.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|