chATLAS_Chains 0.1.3__tar.gz → 0.1.5__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.
Files changed (51) hide show
  1. {chatlas_chains-0.1.3/chATLAS_Chains.egg-info → chatlas_chains-0.1.5}/PKG-INFO +51 -22
  2. chatlas_chains-0.1.5/README.md +116 -0
  3. chatlas_chains-0.1.5/chATLAS_Chains/chains/advanced.py +307 -0
  4. chatlas_chains-0.1.5/chATLAS_Chains/chains/basic.py +89 -0
  5. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/chains/basic_graph.py +40 -19
  6. chatlas_chains-0.1.5/chATLAS_Chains/chains/enhanced_agentic_graph.py +873 -0
  7. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/chains/websearch_retrieval_chain.py +17 -21
  8. chatlas_chains-0.1.5/chATLAS_Chains/documents/rerank.py +96 -0
  9. chatlas_chains-0.1.5/chATLAS_Chains/documents/rrf.py +147 -0
  10. chatlas_chains-0.1.5/chATLAS_Chains/llm/groq.py +598 -0
  11. chatlas_chains-0.1.5/chATLAS_Chains/llm/model_selection.py +172 -0
  12. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/prompt/starters.py +2 -4
  13. chatlas_chains-0.1.5/chATLAS_Chains/query/query_rewriting.py +56 -0
  14. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/search/basic.py +22 -10
  15. chatlas_chains-0.1.5/chATLAS_Chains/utils/doc_utils.py +67 -0
  16. chatlas_chains-0.1.5/chATLAS_Chains/vectorstore.py +190 -0
  17. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5/chATLAS_Chains.egg-info}/PKG-INFO +51 -22
  18. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains.egg-info/SOURCES.txt +7 -3
  19. chatlas_chains-0.1.5/chATLAS_Chains.egg-info/requires.txt +11 -0
  20. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains.egg-info/top_level.txt +0 -1
  21. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/pyproject.toml +21 -5
  22. chatlas_chains-0.1.5/tests/conftest.py +68 -0
  23. chatlas_chains-0.1.5/tests/test_chains.py +155 -0
  24. chatlas_chains-0.1.5/tests/test_groq.py +786 -0
  25. chatlas_chains-0.1.5/tests/test_llm.py +205 -0
  26. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/tests/test_search.py +16 -13
  27. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/tests/test_utils.py +2 -2
  28. chatlas_chains-0.1.3/README.md +0 -92
  29. chatlas_chains-0.1.3/benchmark/basic.py +0 -34
  30. chatlas_chains-0.1.3/benchmark/conversational.py +0 -0
  31. chatlas_chains-0.1.3/chATLAS_Chains/chains/basic.py +0 -42
  32. chatlas_chains-0.1.3/chATLAS_Chains/chains/conversational.py +0 -0
  33. chatlas_chains-0.1.3/chATLAS_Chains/llm/model_selection.py +0 -47
  34. chatlas_chains-0.1.3/chATLAS_Chains/utils/doc_utils.py +0 -28
  35. chatlas_chains-0.1.3/chATLAS_Chains/vectorstore.py +0 -94
  36. chatlas_chains-0.1.3/chATLAS_Chains.egg-info/requires.txt +0 -5
  37. chatlas_chains-0.1.3/tests/conftest.py +0 -274
  38. chatlas_chains-0.1.3/tests/test_chains.py +0 -54
  39. chatlas_chains-0.1.3/tests/test_llm.py +0 -81
  40. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/LICENSE +0 -0
  41. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/__init__.py +0 -0
  42. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/chains/__init__.py +0 -0
  43. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/llm/__init__.py +0 -0
  44. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/log.py +0 -0
  45. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/prompt/__init__.py +0 -0
  46. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  47. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/search/__init__.py +0 -0
  48. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains/utils/__init__.py +0 -0
  49. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  50. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/setup.cfg +0 -0
  51. {chatlas_chains-0.1.3 → chatlas_chains-0.1.5}/tests/__init__.py +0 -0
@@ -1,22 +1,27 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: chATLAS_Chains
3
- Version: 0.1.3
3
+ Version: 0.1.5
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
7
- Project-URL: Homepage, https://gitlab.cern.ch/belliot/chatlas-packages/
8
- Project-URL: Documentation, https://chatlas-packages.docs.cern.ch/chATLAS_Chain/
7
+ Project-URL: Homepage, https://gitlab.cern.ch/atlasml/chatlas/chatlas-packages/
9
8
  Classifier: Programming Language :: Python :: 3
10
9
  Classifier: License :: OSI Approved :: Apache Software License
11
10
  Classifier: Operating System :: OS Independent
12
11
  Requires-Python: >=3.11
13
12
  Description-Content-Type: text/markdown
14
13
  License-File: LICENSE
15
- Requires-Dist: chATLAS_Benchmark>=0.0.9
16
- Requires-Dist: chATLAS_Embed>=0.1.14
14
+ Requires-Dist: chatlas-embed>=0.1.19
17
15
  Requires-Dist: langchain~=0.3.3
18
16
  Requires-Dist: langchain_core
19
17
  Requires-Dist: langchain_openai
18
+ Requires-Dist: langgraph
19
+ Requires-Dist: sentence-transformers>=3.0.0
20
+ Requires-Dist: tiktoken
21
+ Requires-Dist: pinecone
22
+ Requires-Dist: torch==2.2.1
23
+ Requires-Dist: chatlas-embed
24
+ Requires-Dist: psycopg2-binary>=2.9.10
20
25
  Dynamic: license-file
21
26
 
22
27
 
@@ -26,10 +31,18 @@ This package implements and benchmarks various Retrieval Augmented Generation (R
26
31
 
27
32
  ## Installation
28
33
 
34
+ ### From PyPI
35
+
36
+ ```bash
37
+ pip install chATLAS-Chains
38
+ ```
39
+
40
+ ### From source
41
+
42
+ We recommend using [`uv`](https://docs.astral.sh/uv/)
29
43
  ```bash
30
- conda create -n venv chatlas_chains_env python=3.10
31
- conda activate chatlas_chains_env
32
- pip install chatlas-chains
44
+ cd chATLAS_Chains
45
+ uv sync
33
46
  ```
34
47
 
35
48
  ## Environment variables
@@ -38,30 +51,29 @@ These are required for the following use cases
38
51
 
39
52
  1. Using an OpenAI LLM
40
53
  ```bash
41
- export CHATLAS_OPENAI_KEY='your api key'
54
+ export CHATLAS_OPENAI_KEY="your api key"
42
55
  ```
43
56
 
44
- 2. Benchmarking, set the path to the question set
57
+ 2. Using LLMs via the Groq API
45
58
  ```bash
46
- export CHATLAS_BENCHMARK_QUESTIONS=/path/to/questions.josn
59
+ export CHATLAS_GROQ_BASE_URL="http://cs-513-ml003:3000"
60
+ export CHATLAS_GROQ_KEY="your groq api key"
47
61
  ```
48
62
 
49
- ## Available Chains
50
- - chains.basic.basic_retrieval_chain
51
- - chains.basic_graph.basic_retrieval_graph
52
-
53
- ## Benchmarking
54
-
55
- To benchmark e.g. the chains in `chATLAS_Chains.chains.basic` run this from the project root
63
+ **note** The API address is local to the CERN network. If not at CERN, you can forward it like so:
56
64
  ```bash
57
- python benchmark/basic.py
65
+ ssh -L 3000:cs-513-ml003:3000 <LXPLUS_USERNAME>@lxplus.cern.ch
66
+ export CHATLAS_GROQ_BASE_URL="http://localhost:3000"
58
67
  ```
59
68
 
60
- ## Testing
69
+ ## Available Chains
70
+ - chains.basic.basic_retrieval_chain
71
+ - chains.basic_graph.basic_retrieval_graph
72
+ - chains.advanced.advanced_rag
61
73
 
62
- The tests require a running postgres server to work. If on lxplus you can modify `TEST_DB_CONFIG` in [tests/conftest.py](tests/conftest.py) to connect to the chATLAS server.
74
+ ## Postgres
63
75
 
64
- If you want to create a local dummy postgres server, you need to install `psql`. This can be done on macOS using [homebrew](https://brew.sh):
76
+ If you want to create a local postgres server, you need to install `psql`. Some instructions to do this on macOS using [homebrew](https://brew.sh) are here:
65
77
 
66
78
  Software install
67
79
  ```bash
@@ -79,6 +91,23 @@ CREATE EXTENSION IF NOT EXISTS vector;
79
91
  ```
80
92
  ## CHANGELOG
81
93
 
94
+ #### 0.1.5
95
+
96
+ Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
97
+
98
+ #### 0.1.4
99
+
100
+ Support for Groq-hosted models
101
+
102
+ Some new functions that go beyond the "basic RAG" workflow:
103
+ - Reciprocal Rerank Fusion `chATLAS_Chains.documents.rrf.reciprocal_rank_fusion`
104
+ - Document reranking via the Pinecone API `chATLAS_Chains.documents.rerank.rerank_documents`
105
+ - Query rewriting step `chATLAS_Chains.query.query_rewriting.rewrite_query`
106
+
107
+ These are all usable via the new chain `chATLAS_Chains.chains.advanced.advanced_rag`
108
+
109
+ Added unit tests to gitlab CI/CD pipeline
110
+
82
111
  #### 0.1.3
83
112
 
84
113
  Fixing imports
@@ -0,0 +1,116 @@
1
+
2
+ # chATLAS_Chains
3
+
4
+ This package implements and benchmarks various Retrieval Augmented Generation (RAG) chains for use in the [chATLAS](https://chatlas-flask-chatlas.app.cern.ch) project.
5
+
6
+ ## Installation
7
+
8
+ ### From PyPI
9
+
10
+ ```bash
11
+ pip install chATLAS-Chains
12
+ ```
13
+
14
+ ### From source
15
+
16
+ We recommend using [`uv`](https://docs.astral.sh/uv/)
17
+ ```bash
18
+ cd chATLAS_Chains
19
+ uv sync
20
+ ```
21
+
22
+ ## Environment variables
23
+
24
+ These are required for the following use cases
25
+
26
+ 1. Using an OpenAI LLM
27
+ ```bash
28
+ export CHATLAS_OPENAI_KEY="your api key"
29
+ ```
30
+
31
+ 2. Using LLMs via the Groq API
32
+ ```bash
33
+ export CHATLAS_GROQ_BASE_URL="http://cs-513-ml003:3000"
34
+ export CHATLAS_GROQ_KEY="your groq api key"
35
+ ```
36
+
37
+ **note** The API address is local to the CERN network. If not at CERN, you can forward it like so:
38
+ ```bash
39
+ ssh -L 3000:cs-513-ml003:3000 <LXPLUS_USERNAME>@lxplus.cern.ch
40
+ export CHATLAS_GROQ_BASE_URL="http://localhost:3000"
41
+ ```
42
+
43
+ ## Available Chains
44
+ - chains.basic.basic_retrieval_chain
45
+ - chains.basic_graph.basic_retrieval_graph
46
+ - chains.advanced.advanced_rag
47
+
48
+ ## Postgres
49
+
50
+ If you want to create a local postgres server, you need to install `psql`. Some instructions to do this on macOS using [homebrew](https://brew.sh) are here:
51
+
52
+ Software install
53
+ ```bash
54
+ brew install postgresql
55
+ brew services start postgresql
56
+ brew install pgvector
57
+ brew unlink pgvector && brew link pgvector
58
+ ```
59
+
60
+ Create a user
61
+ ```bash
62
+ psql -h localhost -U postgres
63
+ ALTER USER postgres WITH PASSWORD 'Set_your_password_here';
64
+ CREATE EXTENSION IF NOT EXISTS vector;
65
+ ```
66
+ ## CHANGELOG
67
+
68
+ #### 0.1.5
69
+
70
+ Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
71
+
72
+ #### 0.1.4
73
+
74
+ Support for Groq-hosted models
75
+
76
+ Some new functions that go beyond the "basic RAG" workflow:
77
+ - Reciprocal Rerank Fusion `chATLAS_Chains.documents.rrf.reciprocal_rank_fusion`
78
+ - Document reranking via the Pinecone API `chATLAS_Chains.documents.rerank.rerank_documents`
79
+ - Query rewriting step `chATLAS_Chains.query.query_rewriting.rewrite_query`
80
+
81
+ These are all usable via the new chain `chATLAS_Chains.chains.advanced.advanced_rag`
82
+
83
+ Added unit tests to gitlab CI/CD pipeline
84
+
85
+ #### 0.1.3
86
+
87
+ Fixing imports
88
+
89
+ Changed output format of `basic_retrieval_chain` (`docs` key is now a list of `Document` objects, rather than a dict)
90
+
91
+ Unit tests for `basic_retrieval_chain`
92
+
93
+ #### 0.1.2
94
+
95
+ Unit tests
96
+
97
+ First Langgraph chain
98
+
99
+ #### 0.1.1
100
+
101
+ Initial Release
102
+
103
+ ---
104
+ ## 📄 License
105
+
106
+ chATLAS_Benchmark is released under Apache v2.0 license.
107
+
108
+ ---
109
+
110
+ <div align="center">
111
+
112
+ **Made with ❤️ by the ATLAS Collaboration**
113
+
114
+ *For questions and support, please [contact](mailto:joseph.caimin.egan@cern.ch)*
115
+
116
+ </div>
@@ -0,0 +1,307 @@
1
+ """
2
+ More advanced RAG workflow with optional query rewriting, reciprocal rank fusion and reranking
3
+
4
+ Stages:
5
+ - (Optional) Query Rewriting - Correct typos and enhance query clarity using LLM
6
+ - Retrieval - Retrieve documents using BM25 and vector search
7
+ - (Optional) Reciprocal Rank Fusion - Combine results from retrieval modes (e.g. text and vector)
8
+ upweighting results that appear in both
9
+ - (Optional) Reranking - Rerank documents using Pinecone API cross-encoder model
10
+ - Answer Generation - Generate answer using LLM with retrieved context
11
+ """
12
+
13
+ import argparse
14
+ import os
15
+ import sys
16
+ from typing import TypedDict
17
+
18
+ from langchain_core.documents import Document
19
+ from langchain_core.prompts import ChatPromptTemplate
20
+ from langgraph.graph import END, StateGraph
21
+ from langgraph.graph.state import CompiledStateGraph
22
+
23
+ from chATLAS_Chains.documents.rerank import rerank_documents
24
+ from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
25
+ from chATLAS_Chains.llm.groq import RetryConfig
26
+ from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS, get_chat_model
27
+ from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
28
+ from chATLAS_Chains.query.query_rewriting import rewrite_query
29
+ from chATLAS_Chains.search.basic import search_runnable
30
+ from chATLAS_Chains.utils.doc_utils import combine_documents
31
+ from chATLAS_Embed.Base import VectorStore
32
+
33
+
34
+ # Define TypedDict for the simplified state
35
+ class HybridGraphState(TypedDict, total=False):
36
+ question: str
37
+ search_kwargs: dict
38
+ docs: list[Document]
39
+ answer: str
40
+
41
+
42
+ def advanced_rag(
43
+ vectorstore: VectorStore | list[VectorStore],
44
+ model_name: str,
45
+ prompt: str | None = None,
46
+ max_tokens: int | None = None,
47
+ temperature: float = 0.1,
48
+ use_preview_models: bool = False,
49
+ enable_query_rewriting: bool = False,
50
+ enable_rrf: bool = False,
51
+ enable_reranking: bool = False,
52
+ # enable_self_evaluation: bool = False,
53
+ query_rewriting_model: str = GROQ_PRODUCTION_MODELS[0],
54
+ query_rewriting_temperature: float = 0.1,
55
+ rerank_model: str = "cohere-rerank-3.5",
56
+ pinecone_api_key: str | None = None,
57
+ rrf_constant: float = 60.0,
58
+ rrf_weights: dict[str, float] | None = None,
59
+ retry_config: RetryConfig | None = None,
60
+ ) -> CompiledStateGraph:
61
+ """
62
+ Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
63
+
64
+ :param prompt: The prompt template to use for the language model. If None, uses chATLAS_Chains.prompt.starters.CHAT_PROMPT_TEMPLATE
65
+ :param vectorstore: Single vectorstore instance or list of vectorstore instances to search
66
+ :param model_name: The name of the language model to use for generating responses.
67
+ :param max_tokens: Maximum number of tokens to generate in the response. If None, uses the model's default value.
68
+ :param temperature: Temperature to use for the model.
69
+ :param use_preview_models: If True, allows the use of preview models from Groq.
70
+ :param enable_query_rewriting: Whether to enable LLM-powered query rewriting.
71
+ :param enable_rrf: Whether to enable RRF (Reciprocal Rank Fusion) for combining results from multiple vectorstores.
72
+ :param enable_reranking: Whether to rerank the retrieved results using the Pinecone API.
73
+ :param query_rewriting_model: Name of the model to use for query rewriting.
74
+ :param query_rewriting_temperature: Temperature to use for query rewriting.
75
+ :param rerank_model: Name of the Pinecone reranker model to use.
76
+ :param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
77
+ :param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
78
+ :param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
79
+ :param retry_config: Optional RetryConfig for Groq API calls
80
+
81
+ :return: A compiled LangGraph with the chosen features
82
+
83
+ """
84
+ if prompt is None:
85
+ prompt = CHAT_PROMPT_TEMPLATE
86
+
87
+ # Create prompt template
88
+ prompt_template = ChatPromptTemplate.from_template(prompt)
89
+
90
+ # Parallel searcher for vectorstore(s)
91
+ searcher = search_runnable(vectorstore)
92
+
93
+ # Initialise models
94
+ model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
95
+
96
+ if enable_query_rewriting:
97
+ query_rewriting_model_instance = get_chat_model(
98
+ model_name=query_rewriting_model,
99
+ temperature=query_rewriting_temperature,
100
+ use_preview_models=use_preview_models,
101
+ )
102
+
103
+ # --------------- Define functions for the graph nodes ---------------
104
+
105
+ def query_rewriting(state: HybridGraphState) -> HybridGraphState:
106
+ """Rewrite the query to correct typos and enhance clarity."""
107
+ question = state["question"]
108
+
109
+ if not enable_query_rewriting:
110
+ print("[WARNING] Query rewriting disabled, but query_rewriting node was called.")
111
+ # return original state
112
+ return {**state}
113
+
114
+ rewritten_query = rewrite_query(question, model=query_rewriting_model_instance)
115
+
116
+ return {
117
+ **state,
118
+ "question": rewritten_query,
119
+ "unchanged_question": question, # Keep original for reference
120
+ }
121
+
122
+ def retrieval(state: HybridGraphState) -> HybridGraphState:
123
+ """Call the search runnable to retrieve documents"""
124
+
125
+ question = state.get("question")
126
+ if question is None:
127
+ raise Exception("question field is None")
128
+
129
+ search_kwargs = state.get("search_kwargs", {})
130
+ if not search_kwargs:
131
+ print("[WARNING] No search_kwargs provided, using defaults.")
132
+ results = searcher.invoke(question)
133
+ else:
134
+ results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
135
+
136
+ if "docs" not in results:
137
+ raise Exception("Search results missing 'docs' field")
138
+
139
+ print(f"Retrieved {len(results['docs'])} documents")
140
+
141
+ return {
142
+ **state,
143
+ "docs": results["docs"],
144
+ }
145
+
146
+ # Define the document reranking function
147
+ def rerank(state: HybridGraphState) -> HybridGraphState:
148
+ """
149
+ Rerank the parent documents using the Pinecone API.
150
+ """
151
+
152
+ if not enable_reranking:
153
+ print("[WARNING] Reranking disabled, but rerank node was called.")
154
+ # return original state
155
+ return {**state}
156
+
157
+ docs = state.get("docs", [])
158
+
159
+ try:
160
+ reranked_docs = rerank_documents(
161
+ question=state["question"],
162
+ docs=docs,
163
+ reranker_model=rerank_model,
164
+ api_key=pinecone_api_key,
165
+ # num_return_docs = None # return everything
166
+ )
167
+
168
+ except Exception as e:
169
+ print(f"[WARNING] Reranking failed with exception {e}. Returning original documents.")
170
+ reranked_docs = docs
171
+
172
+ return {
173
+ **state,
174
+ "docs": reranked_docs,
175
+ }
176
+
177
+ def rrf(state: HybridGraphState) -> HybridGraphState:
178
+ """Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
179
+ if not enable_rrf:
180
+ print("[WARNING] RRF disabled, but rrf node was called.")
181
+ # return original state
182
+ return {**state}
183
+
184
+ docs = state.get("docs", [])
185
+ if not docs:
186
+ print("[WARNING] No documents retrieved for RRF.")
187
+ return {**state}
188
+
189
+ try:
190
+ rrf_docs = reciprocal_rank_fusion(
191
+ results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
192
+ )
193
+ except Exception as e:
194
+ print(f"[WARNING] RRF failed with exception {e}. Returning original documents.")
195
+ rrf_docs = docs
196
+
197
+ return {
198
+ **state,
199
+ "docs": rrf_docs,
200
+ }
201
+
202
+ def generate_answer(state: HybridGraphState) -> HybridGraphState:
203
+ """Generate a response from the LLM"""
204
+
205
+ # Format the prompt using LangChain template
206
+ prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
207
+ final_prompt = prompt_template.format_messages(**prompt_input)
208
+
209
+ response = model.invoke(final_prompt)
210
+ answer = response.content
211
+
212
+ return {
213
+ **state,
214
+ "answer": answer,
215
+ }
216
+
217
+ # --------------- Build the graph ---------------
218
+ graph = StateGraph(HybridGraphState)
219
+
220
+ # Add all the nodes, but don't link to them if not using
221
+ graph.add_node("query_rewrite", query_rewriting)
222
+ graph.add_node("retrieval", retrieval)
223
+ graph.add_node("rrf", rrf)
224
+ graph.add_node("rerank", rerank)
225
+ graph.add_node("generate", generate_answer)
226
+ # graph.add_node("assess", assess_answer)
227
+ # graph.add_node("refine", refine_answer)
228
+
229
+ if enable_query_rewriting:
230
+ # rewrite the query first
231
+ graph.add_edge("query_rewrite", "retrieval")
232
+ graph.set_entry_point("query_rewrite")
233
+ else:
234
+ # start with retrieval
235
+ graph.set_entry_point("retrieval")
236
+
237
+ if not enable_rrf and not enable_reranking:
238
+ # no document processing, go straight to generation
239
+ graph.add_edge("retrieval", "generate")
240
+
241
+ elif enable_rrf and not enable_reranking:
242
+ # retrieve → rrf → generate
243
+ graph.add_edge("retrieval", "rrf")
244
+ graph.add_edge("rrf", "generate")
245
+
246
+ elif not enable_rrf and enable_reranking:
247
+ # retrieve → rerank → generate
248
+ graph.add_edge("retrieval", "rerank")
249
+ graph.add_edge("rerank", "generate")
250
+
251
+ else:
252
+ # retrieve → rrf → rerank → generate
253
+ graph.add_edge("retrieval", "rrf")
254
+ graph.add_edge("rrf", "rerank")
255
+ graph.add_edge("rerank", "generate")
256
+
257
+ # generate at the end
258
+ graph.add_edge("generate", END)
259
+
260
+ # Compile the graph
261
+ return graph.compile()
262
+
263
+
264
+ if __name__ == "__main__":
265
+ from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
266
+ from chATLAS_Chains.vectorstore import get_vectorstore
267
+
268
+ twiki = get_vectorstore("twiki_prod")
269
+
270
+ retry_config = RetryConfig(
271
+ max_retries=1,
272
+ max_delay=120.0,
273
+ )
274
+
275
+ # Create the hybrid graph
276
+ graph = advanced_rag(
277
+ vectorstore=[twiki],
278
+ model_name=GROQ_PRODUCTION_MODELS[0],
279
+ enable_query_rewriting=True,
280
+ enable_rrf=True,
281
+ enable_reranking=True,
282
+ )
283
+
284
+ # Test query
285
+ try:
286
+ ans = graph.invoke(
287
+ {
288
+ "question": "How can one check for and remove bad or corrupted events in the analysis?",
289
+ "search_kwargs": {
290
+ "k_text": 3,
291
+ "k": 15,
292
+ "date_filter": "01-01-2010",
293
+ # "type": ["twiki"],
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
305
+
306
+ traceback.print_exc()
307
+ sys.exit(1)
@@ -0,0 +1,89 @@
1
+ from operator import itemgetter
2
+ from typing import Optional
3
+
4
+ from langchain_core.prompts import ChatPromptTemplate
5
+ from langchain_core.runnables import RunnableSerializable
6
+
7
+ from chATLAS_Chains.llm.groq import RetryConfig
8
+ from chATLAS_Chains.llm.model_selection import get_chat_model
9
+ from chATLAS_Chains.search.basic import search_runnable
10
+ from chATLAS_Chains.utils.doc_utils import combine_documents
11
+ from chATLAS_Embed.Base import VectorStore
12
+
13
+
14
+ def basic_retrieval_chain(
15
+ prompt: str,
16
+ vectorstore: VectorStore | list[VectorStore],
17
+ model_name: str,
18
+ max_tokens: int | None = None,
19
+ temperature: float | None = None,
20
+ use_preview_models: bool = False,
21
+ retry_config: RetryConfig | None = None,
22
+ ) -> RunnableSerializable:
23
+ """
24
+ Baseline RAG retrieval chain. Searches one or several vectorstores in parallel, passes retrieved documents to the model
25
+
26
+ :param prompt: The prompt template to use with the model.
27
+ :type prompt: str
28
+ :param vectorstore: The vectorstore or list of vectorstores to search over.
29
+ :type vectorstore: Any
30
+ :param model_name: The name of the chat model to use for generating responses.
31
+ :type model_name: str
32
+ :param max_tokens: The maximum number of tokens to generate in the response. Defaults to None.
33
+ :type max_tokens: int | None
34
+ :param temperature: The temperature to use for the model's response generation. Defaults to None.
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
40
+
41
+ :return: A LangChain RunnableSerializable chain that performs retrieval and response generation.
42
+ :rtype: RunnableSerializable
43
+
44
+ """
45
+ prompt_template = ChatPromptTemplate.from_template(prompt)
46
+ model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
47
+
48
+ search = search_runnable(vectorstore)
49
+
50
+ final_inputs = {
51
+ "context": lambda x: combine_documents(x["docs"]),
52
+ "question": itemgetter("question"),
53
+ }
54
+
55
+ answer = {
56
+ "answer": final_inputs | prompt_template | model,
57
+ "docs": lambda x: x["docs"],
58
+ }
59
+
60
+ chain = search | answer
61
+ return chain
62
+
63
+
64
+ if __name__ == "__main__":
65
+ from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
66
+ from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
67
+ from chATLAS_Chains.vectorstore import get_vectorstore
68
+
69
+ twiki_vectorstore = get_vectorstore("twiki_prod")
70
+ mkdocs_vectorstore = get_vectorstore("mkdocs_prod_v1")
71
+
72
+ chain = basic_retrieval_chain(
73
+ prompt=CHAT_PROMPT_TEMPLATE,
74
+ vectorstore=[twiki_vectorstore, mkdocs_vectorstore],
75
+ model_name=GROQ_PRODUCTION_MODELS[0],
76
+ # model_name="meta-llama/llama-4-maverick-17b-128e-instruct",
77
+ # model_name="gemma2-9b-it",
78
+ # model_name="qwen-qwq-32b",
79
+ # model_name="mistral-saba-24b",
80
+ )
81
+ SEARCH_HYPERPARAMS = {"k": 5, "k_text": 0, "date_filter": "01-01-2010"}
82
+
83
+ result = chain.invoke("What is the Higgs boson?", config={"metadata": {"search_kwargs": SEARCH_HYPERPARAMS}})
84
+
85
+ print(f"Answer: {result['answer'].content}")
86
+ print(f"Number of documents retrieved: {len(result['docs'])}")
87
+
88
+ for doc in result["docs"]:
89
+ print(doc.metadata.get("source"))