chATLAS_Chains 0.1.4__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.4 → chatlas_chains-0.1.6}/PKG-INFO +11 -1
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/README.md +10 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/advanced.py +41 -36
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic.py +8 -1
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rrf.py +29 -26
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/groq.py +2 -1
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/model_selection.py +11 -7
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/PKG-INFO +11 -1
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/SOURCES.txt +1 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/pyproject.toml +1 -1
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_chains.py +58 -31
- {chatlas_chains-0.1.4 → 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.4 → chatlas_chains-0.1.6}/LICENSE +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/__init__.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/__init__.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic_graph.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rerank.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/__init__.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/log.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/__init__.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/starters.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/query/query_rewriting.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/search/__init__.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/search/basic.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/__init__.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/doc_utils.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/vectorstore.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/requires.txt +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/top_level.txt +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/setup.cfg +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/__init__.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/conftest.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_groq.py +0 -0
- {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_search.py +0 -0
- {chatlas_chains-0.1.4 → 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,16 @@ 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
|
+
|
|
100
|
+
#### 0.1.5
|
|
101
|
+
|
|
102
|
+
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
103
|
+
|
|
94
104
|
#### 0.1.4
|
|
95
105
|
|
|
96
106
|
Support for Groq-hosted models
|
|
@@ -65,6 +65,16 @@ 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
|
+
|
|
74
|
+
#### 0.1.5
|
|
75
|
+
|
|
76
|
+
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
77
|
+
|
|
68
78
|
#### 0.1.4
|
|
69
79
|
|
|
70
80
|
Support for Groq-hosted models
|
|
@@ -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
|
|
@@ -22,6 +23,7 @@ from langgraph.graph.state import CompiledStateGraph
|
|
|
22
23
|
|
|
23
24
|
from chATLAS_Chains.documents.rerank import rerank_documents
|
|
24
25
|
from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
|
|
26
|
+
from chATLAS_Chains.llm.groq import RetryConfig
|
|
25
27
|
from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS, get_chat_model
|
|
26
28
|
from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
|
|
27
29
|
from chATLAS_Chains.query.query_rewriting import rewrite_query
|
|
@@ -29,6 +31,8 @@ from chATLAS_Chains.search.basic import search_runnable
|
|
|
29
31
|
from chATLAS_Chains.utils.doc_utils import combine_documents
|
|
30
32
|
from chATLAS_Embed.Base import VectorStore
|
|
31
33
|
|
|
34
|
+
logger = logging.getLogger(__name__)
|
|
35
|
+
|
|
32
36
|
|
|
33
37
|
# Define TypedDict for the simplified state
|
|
34
38
|
class HybridGraphState(TypedDict, total=False):
|
|
@@ -55,6 +59,8 @@ def advanced_rag(
|
|
|
55
59
|
pinecone_api_key: str | None = None,
|
|
56
60
|
rrf_constant: float = 60.0,
|
|
57
61
|
rrf_weights: dict[str, float] | None = None,
|
|
62
|
+
retry_config: RetryConfig | None = None,
|
|
63
|
+
fallback_models: list[str] | None = None,
|
|
58
64
|
) -> CompiledStateGraph:
|
|
59
65
|
"""
|
|
60
66
|
Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
|
|
@@ -73,7 +79,9 @@ def advanced_rag(
|
|
|
73
79
|
:param rerank_model: Name of the Pinecone reranker model to use.
|
|
74
80
|
:param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
|
|
75
81
|
:param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
|
|
76
|
-
:rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
|
|
82
|
+
:param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
|
|
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.
|
|
77
85
|
|
|
78
86
|
:return: A compiled LangGraph with the chosen features
|
|
79
87
|
|
|
@@ -88,13 +96,15 @@ def advanced_rag(
|
|
|
88
96
|
searcher = search_runnable(vectorstore)
|
|
89
97
|
|
|
90
98
|
# Initialise models
|
|
91
|
-
model = get_chat_model(model_name, max_tokens, temperature, use_preview_models)
|
|
99
|
+
model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config, fallback_models)
|
|
92
100
|
|
|
93
101
|
if enable_query_rewriting:
|
|
94
102
|
query_rewriting_model_instance = get_chat_model(
|
|
95
103
|
model_name=query_rewriting_model,
|
|
96
104
|
temperature=query_rewriting_temperature,
|
|
97
105
|
use_preview_models=use_preview_models,
|
|
106
|
+
retry_config=retry_config,
|
|
107
|
+
fallback_models=fallback_models,
|
|
98
108
|
)
|
|
99
109
|
|
|
100
110
|
# --------------- Define functions for the graph nodes ---------------
|
|
@@ -104,7 +114,7 @@ def advanced_rag(
|
|
|
104
114
|
question = state["question"]
|
|
105
115
|
|
|
106
116
|
if not enable_query_rewriting:
|
|
107
|
-
|
|
117
|
+
logger.warning("Query rewriting disabled, but query_rewriting node was called.")
|
|
108
118
|
# return original state
|
|
109
119
|
return {**state}
|
|
110
120
|
|
|
@@ -125,7 +135,7 @@ def advanced_rag(
|
|
|
125
135
|
|
|
126
136
|
search_kwargs = state.get("search_kwargs", {})
|
|
127
137
|
if not search_kwargs:
|
|
128
|
-
|
|
138
|
+
logger.warning("No search_kwargs provided, using defaults.")
|
|
129
139
|
results = searcher.invoke(question)
|
|
130
140
|
else:
|
|
131
141
|
results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
|
|
@@ -133,7 +143,7 @@ def advanced_rag(
|
|
|
133
143
|
if "docs" not in results:
|
|
134
144
|
raise Exception("Search results missing 'docs' field")
|
|
135
145
|
|
|
136
|
-
|
|
146
|
+
logger.debug(f"Retrieved {len(results['docs'])} documents")
|
|
137
147
|
|
|
138
148
|
return {
|
|
139
149
|
**state,
|
|
@@ -147,7 +157,7 @@ def advanced_rag(
|
|
|
147
157
|
"""
|
|
148
158
|
|
|
149
159
|
if not enable_reranking:
|
|
150
|
-
|
|
160
|
+
logger.warning("Reranking disabled, but rerank node was called.")
|
|
151
161
|
# return original state
|
|
152
162
|
return {**state}
|
|
153
163
|
|
|
@@ -163,7 +173,7 @@ def advanced_rag(
|
|
|
163
173
|
)
|
|
164
174
|
|
|
165
175
|
except Exception as e:
|
|
166
|
-
|
|
176
|
+
logger.warning(f"Reranking failed with exception {e}. Returning original documents.")
|
|
167
177
|
reranked_docs = docs
|
|
168
178
|
|
|
169
179
|
return {
|
|
@@ -174,13 +184,13 @@ def advanced_rag(
|
|
|
174
184
|
def rrf(state: HybridGraphState) -> HybridGraphState:
|
|
175
185
|
"""Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
|
|
176
186
|
if not enable_rrf:
|
|
177
|
-
|
|
187
|
+
logger.warning("RRF disabled, but rrf node was called.")
|
|
178
188
|
# return original state
|
|
179
189
|
return {**state}
|
|
180
190
|
|
|
181
191
|
docs = state.get("docs", [])
|
|
182
192
|
if not docs:
|
|
183
|
-
|
|
193
|
+
logger.warning("No documents retrieved for RRF.")
|
|
184
194
|
return {**state}
|
|
185
195
|
|
|
186
196
|
try:
|
|
@@ -188,7 +198,7 @@ def advanced_rag(
|
|
|
188
198
|
results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
|
|
189
199
|
)
|
|
190
200
|
except Exception as e:
|
|
191
|
-
|
|
201
|
+
logger.warning(f"RRF failed with exception {e}. Returning original documents.")
|
|
192
202
|
rrf_docs = docs
|
|
193
203
|
|
|
194
204
|
return {
|
|
@@ -198,7 +208,6 @@ def advanced_rag(
|
|
|
198
208
|
|
|
199
209
|
def generate_answer(state: HybridGraphState) -> HybridGraphState:
|
|
200
210
|
"""Generate a response from the LLM"""
|
|
201
|
-
|
|
202
211
|
# Format the prompt using LangChain template
|
|
203
212
|
prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
|
|
204
213
|
final_prompt = prompt_template.format_messages(**prompt_input)
|
|
@@ -263,37 +272,33 @@ if __name__ == "__main__":
|
|
|
263
272
|
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
264
273
|
|
|
265
274
|
twiki = get_vectorstore("twiki_prod")
|
|
275
|
+
mkdocs = get_vectorstore("mkdocs_prod_v1")
|
|
276
|
+
|
|
277
|
+
retry_config = RetryConfig(
|
|
278
|
+
max_retries=1,
|
|
279
|
+
max_delay=120.0,
|
|
280
|
+
)
|
|
266
281
|
|
|
267
282
|
# Create the hybrid graph
|
|
268
283
|
graph = advanced_rag(
|
|
269
|
-
vectorstore=[twiki],
|
|
284
|
+
vectorstore=[twiki, mkdocs],
|
|
270
285
|
model_name=GROQ_PRODUCTION_MODELS[0],
|
|
271
286
|
enable_query_rewriting=True,
|
|
272
287
|
enable_rrf=True,
|
|
273
|
-
enable_reranking=
|
|
288
|
+
enable_reranking=False,
|
|
274
289
|
)
|
|
275
290
|
|
|
276
|
-
|
|
277
|
-
|
|
278
|
-
|
|
279
|
-
{
|
|
280
|
-
"
|
|
281
|
-
"
|
|
282
|
-
|
|
283
|
-
|
|
284
|
-
|
|
285
|
-
|
|
286
|
-
|
|
287
|
-
}
|
|
288
|
-
)
|
|
289
|
-
|
|
290
|
-
print(f"Number of docs is : {len(ans['docs'])}")
|
|
291
|
-
print(f"Answer: {ans['answer']}")
|
|
292
|
-
|
|
293
|
-
except Exception as e:
|
|
294
|
-
print(f"❌ Graph execution failed with error: {e}")
|
|
295
|
-
print(f"Error type: {type(e).__name__}")
|
|
296
|
-
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
|
+
)
|
|
297
302
|
|
|
298
|
-
|
|
299
|
-
|
|
303
|
+
print(f"Number of docs is : {len(ans['docs'])}")
|
|
304
|
+
print(f"Answer: {ans['answer']}")
|
|
@@ -4,6 +4,7 @@ from typing import Optional
|
|
|
4
4
|
from langchain_core.prompts import ChatPromptTemplate
|
|
5
5
|
from langchain_core.runnables import RunnableSerializable
|
|
6
6
|
|
|
7
|
+
from chATLAS_Chains.llm.groq import RetryConfig
|
|
7
8
|
from chATLAS_Chains.llm.model_selection import get_chat_model
|
|
8
9
|
from chATLAS_Chains.search.basic import search_runnable
|
|
9
10
|
from chATLAS_Chains.utils.doc_utils import combine_documents
|
|
@@ -16,6 +17,8 @@ def basic_retrieval_chain(
|
|
|
16
17
|
model_name: str,
|
|
17
18
|
max_tokens: int | None = None,
|
|
18
19
|
temperature: float | None = None,
|
|
20
|
+
use_preview_models: bool = False,
|
|
21
|
+
retry_config: RetryConfig | None = None,
|
|
19
22
|
) -> RunnableSerializable:
|
|
20
23
|
"""
|
|
21
24
|
Baseline RAG retrieval chain. Searches one or several vectorstores in parallel, passes retrieved documents to the model
|
|
@@ -30,13 +33,17 @@ def basic_retrieval_chain(
|
|
|
30
33
|
:type max_tokens: int | None
|
|
31
34
|
:param temperature: The temperature to use for the model's response generation. Defaults to None.
|
|
32
35
|
:type temperature: float | None
|
|
36
|
+
:param use_preview_models: Whether to allow the use of preview models from Groq. Defaults to False.
|
|
37
|
+
:type use_preview_models: bool
|
|
38
|
+
:param retry_config: Configuration for retrying requests to the model in case of failures. Defaults to None, which will use the default retry configuration.
|
|
39
|
+
:type retry_config: RetryConfig | None
|
|
33
40
|
|
|
34
41
|
:return: A LangChain RunnableSerializable chain that performs retrieval and response generation.
|
|
35
42
|
:rtype: RunnableSerializable
|
|
36
43
|
|
|
37
44
|
"""
|
|
38
45
|
prompt_template = ChatPromptTemplate.from_template(prompt)
|
|
39
|
-
model = get_chat_model(model_name, max_tokens, temperature)
|
|
46
|
+
model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
|
|
40
47
|
|
|
41
48
|
search = search_runnable(vectorstore)
|
|
42
49
|
|
|
@@ -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()
|
|
@@ -26,8 +26,8 @@ GROQ_PREVIEW_MODELS = [
|
|
|
26
26
|
"meta-llama/llama-prompt-guard-2-22m",
|
|
27
27
|
"meta-llama/llama-prompt-guard-2-86m",
|
|
28
28
|
"moonshotai/kimi-k2-instruct",
|
|
29
|
-
# "playai-tts",
|
|
30
|
-
# "playai-tts-arabic",
|
|
29
|
+
# "playai-tts", # text to audio model
|
|
30
|
+
# "playai-tts-arabic", # text to audio model
|
|
31
31
|
"qwen/qwen3-32b",
|
|
32
32
|
"gemma2-9b-it",
|
|
33
33
|
]
|
|
@@ -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,16 @@ 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
|
+
|
|
100
|
+
#### 0.1.5
|
|
101
|
+
|
|
102
|
+
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
103
|
+
|
|
94
104
|
#### 0.1.4
|
|
95
105
|
|
|
96
106
|
Support for Groq-hosted models
|
|
@@ -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"
|
|
@@ -13,16 +13,19 @@ from pydantic import ValidationError
|
|
|
13
13
|
|
|
14
14
|
from chATLAS_Chains.chains.advanced import advanced_rag
|
|
15
15
|
from chATLAS_Chains.chains.basic import basic_retrieval_chain
|
|
16
|
+
from chATLAS_Chains.llm.groq import RetryConfig
|
|
16
17
|
from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
|
|
17
18
|
from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
|
|
18
19
|
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
19
20
|
|
|
20
21
|
|
|
21
|
-
|
|
22
|
-
|
|
23
|
-
|
|
24
|
-
|
|
22
|
+
@pytest.fixture
|
|
23
|
+
def retry_config():
|
|
24
|
+
"""Retry config to use in tests"""
|
|
25
|
+
return RetryConfig(max_retries=5, max_delay=120.0)
|
|
26
|
+
|
|
25
27
|
|
|
28
|
+
def test_basic_retrieval_chain_returns_runnablesequence(twiki_vectorstore):
|
|
26
29
|
chain = basic_retrieval_chain(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=twiki_vectorstore, model_name="gpt-4o-mini")
|
|
27
30
|
|
|
28
31
|
assert isinstance(chain, RunnableSequence)
|
|
@@ -30,7 +33,7 @@ def test_basic_retrieval_chain_returns_runnablesequence(twiki_vectorstore):
|
|
|
30
33
|
|
|
31
34
|
def test_search_runnable_returns_runnablesequence_with_list_of_vectorstores(three_vectorstores):
|
|
32
35
|
chain = basic_retrieval_chain(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=three_vectorstores, model_name="gpt-4o-mini")
|
|
33
|
-
|
|
36
|
+
|
|
34
37
|
assert isinstance(chain, RunnableSequence)
|
|
35
38
|
|
|
36
39
|
|
|
@@ -43,28 +46,36 @@ def test_search_runnable_error_on_invalid_input():
|
|
|
43
46
|
)
|
|
44
47
|
|
|
45
48
|
|
|
46
|
-
def test_basic_retrieval_chain_returns_docs_and_answer(twiki_vectorstore):
|
|
49
|
+
def test_basic_retrieval_chain_returns_docs_and_answer(twiki_vectorstore, retry_config):
|
|
47
50
|
chain = basic_retrieval_chain(
|
|
48
|
-
prompt=CHAT_PROMPT_TEMPLATE,
|
|
51
|
+
prompt=CHAT_PROMPT_TEMPLATE,
|
|
52
|
+
vectorstore=twiki_vectorstore,
|
|
53
|
+
model_name=GROQ_PRODUCTION_MODELS[0],
|
|
54
|
+
retry_config=retry_config,
|
|
49
55
|
)
|
|
50
56
|
output = chain.invoke("What is the Higgs boson?")
|
|
51
57
|
|
|
52
|
-
assert "docs" in output
|
|
53
|
-
assert "answer" in output
|
|
58
|
+
assert "docs" in output, "'docs' key not in output"
|
|
59
|
+
assert "answer" in output, "'answer' key not in output"
|
|
54
60
|
|
|
55
|
-
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"
|
|
56
63
|
|
|
57
64
|
|
|
58
|
-
def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores):
|
|
65
|
+
def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores, retry_config):
|
|
59
66
|
chain = basic_retrieval_chain(
|
|
60
|
-
prompt=CHAT_PROMPT_TEMPLATE,
|
|
67
|
+
prompt=CHAT_PROMPT_TEMPLATE,
|
|
68
|
+
vectorstore=three_vectorstores,
|
|
69
|
+
model_name=GROQ_PRODUCTION_MODELS[0],
|
|
70
|
+
retry_config=retry_config,
|
|
61
71
|
)
|
|
62
72
|
output = chain.invoke("What is the Higgs boson?")
|
|
63
73
|
|
|
64
|
-
assert "docs" in output
|
|
65
|
-
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"
|
|
66
77
|
|
|
67
|
-
assert isinstance(output["docs"], list)
|
|
78
|
+
assert isinstance(output["docs"], list), "'docs' is not a list"
|
|
68
79
|
|
|
69
80
|
sources = set([doc.metadata.get("source") for doc in output["docs"]])
|
|
70
81
|
expected_sources = {"twiki", "MkDocs", "CDS"}
|
|
@@ -73,7 +84,7 @@ def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores):
|
|
|
73
84
|
)
|
|
74
85
|
|
|
75
86
|
|
|
76
|
-
def test_advanced_rag_chain(twiki_vectorstore):
|
|
87
|
+
def test_advanced_rag_chain(twiki_vectorstore, retry_config):
|
|
77
88
|
# TODO: this is required currently requried for search to work
|
|
78
89
|
search_kwargs = {
|
|
79
90
|
"k_text": 3,
|
|
@@ -83,16 +94,22 @@ def test_advanced_rag_chain(twiki_vectorstore):
|
|
|
83
94
|
}
|
|
84
95
|
|
|
85
96
|
chain = advanced_rag(
|
|
86
|
-
prompt=CHAT_PROMPT_TEMPLATE,
|
|
97
|
+
prompt=CHAT_PROMPT_TEMPLATE,
|
|
98
|
+
vectorstore=twiki_vectorstore,
|
|
99
|
+
model_name=GROQ_PRODUCTION_MODELS[0],
|
|
100
|
+
retry_config=retry_config,
|
|
87
101
|
)
|
|
88
102
|
output = chain.invoke(
|
|
89
103
|
{"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
|
|
90
104
|
)
|
|
91
105
|
|
|
92
|
-
assert "docs" in output
|
|
93
|
-
assert "
|
|
94
|
-
assert
|
|
95
|
-
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
|
+
)
|
|
96
113
|
|
|
97
114
|
# test with the additional options turned on (besides rerank)
|
|
98
115
|
chain = advanced_rag(
|
|
@@ -107,13 +124,17 @@ def test_advanced_rag_chain(twiki_vectorstore):
|
|
|
107
124
|
{"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
|
|
108
125
|
)
|
|
109
126
|
|
|
110
|
-
assert "docs" in output
|
|
111
|
-
assert "
|
|
112
|
-
assert
|
|
113
|
-
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"
|
|
114
135
|
|
|
115
136
|
|
|
116
|
-
def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores):
|
|
137
|
+
def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores, retry_config):
|
|
117
138
|
search_kwargs = {
|
|
118
139
|
"k_text": 3,
|
|
119
140
|
"k": 15,
|
|
@@ -122,16 +143,22 @@ def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores):
|
|
|
122
143
|
}
|
|
123
144
|
|
|
124
145
|
chain = advanced_rag(
|
|
125
|
-
prompt=CHAT_PROMPT_TEMPLATE,
|
|
146
|
+
prompt=CHAT_PROMPT_TEMPLATE,
|
|
147
|
+
vectorstore=three_vectorstores,
|
|
148
|
+
model_name=GROQ_PRODUCTION_MODELS[0],
|
|
149
|
+
retry_config=retry_config,
|
|
126
150
|
)
|
|
127
151
|
output = chain.invoke(
|
|
128
152
|
{"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
|
|
129
153
|
)
|
|
130
154
|
|
|
131
|
-
assert "docs" in output
|
|
132
|
-
assert "
|
|
133
|
-
assert
|
|
134
|
-
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
|
+
)
|
|
135
162
|
|
|
136
163
|
sources = set([doc.metadata.get("source") for doc in output["docs"]])
|
|
137
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
|
{chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/enhanced_agentic_graph.py
RENAMED
|
File without changes
|
{chatlas_chains-0.1.4 → 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
|