chATLAS_Chains 0.1.3__tar.gz → 0.1.4__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 (50) hide show
  1. {chatlas_chains-0.1.3/chATLAS_Chains.egg-info → chatlas_chains-0.1.4}/PKG-INFO +47 -22
  2. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/README.md +37 -17
  3. chatlas_chains-0.1.4/chATLAS_Chains/chains/advanced.py +299 -0
  4. chatlas_chains-0.1.4/chATLAS_Chains/chains/basic.py +82 -0
  5. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/chains/basic_graph.py +40 -19
  6. chatlas_chains-0.1.4/chATLAS_Chains/chains/enhanced_agentic_graph.py +873 -0
  7. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/chains/websearch_retrieval_chain.py +17 -21
  8. chatlas_chains-0.1.4/chATLAS_Chains/documents/rerank.py +96 -0
  9. chatlas_chains-0.1.4/chATLAS_Chains/documents/rrf.py +147 -0
  10. chatlas_chains-0.1.4/chATLAS_Chains/llm/groq.py +598 -0
  11. chatlas_chains-0.1.4/chATLAS_Chains/llm/model_selection.py +172 -0
  12. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/prompt/starters.py +2 -4
  13. chatlas_chains-0.1.4/chATLAS_Chains/query/query_rewriting.py +56 -0
  14. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/search/basic.py +22 -10
  15. chatlas_chains-0.1.4/chATLAS_Chains/utils/doc_utils.py +67 -0
  16. chatlas_chains-0.1.4/chATLAS_Chains/vectorstore.py +190 -0
  17. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4/chATLAS_Chains.egg-info}/PKG-INFO +47 -22
  18. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains.egg-info/SOURCES.txt +7 -3
  19. chatlas_chains-0.1.4/chATLAS_Chains.egg-info/requires.txt +11 -0
  20. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains.egg-info/top_level.txt +0 -1
  21. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/pyproject.toml +21 -5
  22. chatlas_chains-0.1.4/tests/conftest.py +68 -0
  23. chatlas_chains-0.1.4/tests/test_chains.py +140 -0
  24. chatlas_chains-0.1.4/tests/test_groq.py +786 -0
  25. chatlas_chains-0.1.4/tests/test_llm.py +205 -0
  26. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/tests/test_search.py +16 -13
  27. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/tests/test_utils.py +2 -2
  28. chatlas_chains-0.1.3/benchmark/basic.py +0 -34
  29. chatlas_chains-0.1.3/benchmark/conversational.py +0 -0
  30. chatlas_chains-0.1.3/chATLAS_Chains/chains/basic.py +0 -42
  31. chatlas_chains-0.1.3/chATLAS_Chains/chains/conversational.py +0 -0
  32. chatlas_chains-0.1.3/chATLAS_Chains/llm/model_selection.py +0 -47
  33. chatlas_chains-0.1.3/chATLAS_Chains/utils/doc_utils.py +0 -28
  34. chatlas_chains-0.1.3/chATLAS_Chains/vectorstore.py +0 -94
  35. chatlas_chains-0.1.3/chATLAS_Chains.egg-info/requires.txt +0 -5
  36. chatlas_chains-0.1.3/tests/conftest.py +0 -274
  37. chatlas_chains-0.1.3/tests/test_chains.py +0 -54
  38. chatlas_chains-0.1.3/tests/test_llm.py +0 -81
  39. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/LICENSE +0 -0
  40. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/__init__.py +0 -0
  41. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/chains/__init__.py +0 -0
  42. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/llm/__init__.py +0 -0
  43. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/log.py +0 -0
  44. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/prompt/__init__.py +0 -0
  45. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  46. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/search/__init__.py +0 -0
  47. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/utils/__init__.py +0 -0
  48. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  49. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/setup.cfg +0 -0
  50. {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/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.4
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,19 @@ CREATE EXTENSION IF NOT EXISTS vector;
79
91
  ```
80
92
  ## CHANGELOG
81
93
 
94
+ #### 0.1.4
95
+
96
+ Support for Groq-hosted models
97
+
98
+ Some new functions that go beyond the "basic RAG" workflow:
99
+ - Reciprocal Rerank Fusion `chATLAS_Chains.documents.rrf.reciprocal_rank_fusion`
100
+ - Document reranking via the Pinecone API `chATLAS_Chains.documents.rerank.rerank_documents`
101
+ - Query rewriting step `chATLAS_Chains.query.query_rewriting.rewrite_query`
102
+
103
+ These are all usable via the new chain `chATLAS_Chains.chains.advanced.advanced_rag`
104
+
105
+ Added unit tests to gitlab CI/CD pipeline
106
+
82
107
  #### 0.1.3
83
108
 
84
109
  Fixing imports
@@ -5,10 +5,18 @@ This package implements and benchmarks various Retrieval Augmented Generation (R
5
5
 
6
6
  ## Installation
7
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/)
8
17
  ```bash
9
- conda create -n venv chatlas_chains_env python=3.10
10
- conda activate chatlas_chains_env
11
- pip install chatlas-chains
18
+ cd chATLAS_Chains
19
+ uv sync
12
20
  ```
13
21
 
14
22
  ## Environment variables
@@ -17,30 +25,29 @@ These are required for the following use cases
17
25
 
18
26
  1. Using an OpenAI LLM
19
27
  ```bash
20
- export CHATLAS_OPENAI_KEY='your api key'
28
+ export CHATLAS_OPENAI_KEY="your api key"
21
29
  ```
22
30
 
23
- 2. Benchmarking, set the path to the question set
31
+ 2. Using LLMs via the Groq API
24
32
  ```bash
25
- export CHATLAS_BENCHMARK_QUESTIONS=/path/to/questions.josn
33
+ export CHATLAS_GROQ_BASE_URL="http://cs-513-ml003:3000"
34
+ export CHATLAS_GROQ_KEY="your groq api key"
26
35
  ```
27
36
 
28
- ## Available Chains
29
- - chains.basic.basic_retrieval_chain
30
- - chains.basic_graph.basic_retrieval_graph
31
-
32
- ## Benchmarking
33
-
34
- To benchmark e.g. the chains in `chATLAS_Chains.chains.basic` run this from the project root
37
+ **note** The API address is local to the CERN network. If not at CERN, you can forward it like so:
35
38
  ```bash
36
- python benchmark/basic.py
39
+ ssh -L 3000:cs-513-ml003:3000 <LXPLUS_USERNAME>@lxplus.cern.ch
40
+ export CHATLAS_GROQ_BASE_URL="http://localhost:3000"
37
41
  ```
38
42
 
39
- ## Testing
43
+ ## Available Chains
44
+ - chains.basic.basic_retrieval_chain
45
+ - chains.basic_graph.basic_retrieval_graph
46
+ - chains.advanced.advanced_rag
40
47
 
41
- 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.
48
+ ## Postgres
42
49
 
43
- 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):
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:
44
51
 
45
52
  Software install
46
53
  ```bash
@@ -58,6 +65,19 @@ CREATE EXTENSION IF NOT EXISTS vector;
58
65
  ```
59
66
  ## CHANGELOG
60
67
 
68
+ #### 0.1.4
69
+
70
+ Support for Groq-hosted models
71
+
72
+ Some new functions that go beyond the "basic RAG" workflow:
73
+ - Reciprocal Rerank Fusion `chATLAS_Chains.documents.rrf.reciprocal_rank_fusion`
74
+ - Document reranking via the Pinecone API `chATLAS_Chains.documents.rerank.rerank_documents`
75
+ - Query rewriting step `chATLAS_Chains.query.query_rewriting.rewrite_query`
76
+
77
+ These are all usable via the new chain `chATLAS_Chains.chains.advanced.advanced_rag`
78
+
79
+ Added unit tests to gitlab CI/CD pipeline
80
+
61
81
  #### 0.1.3
62
82
 
63
83
  Fixing imports
@@ -0,0 +1,299 @@
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.model_selection import GROQ_PRODUCTION_MODELS, get_chat_model
26
+ from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
27
+ from chATLAS_Chains.query.query_rewriting import rewrite_query
28
+ from chATLAS_Chains.search.basic import search_runnable
29
+ from chATLAS_Chains.utils.doc_utils import combine_documents
30
+ from chATLAS_Embed.Base import VectorStore
31
+
32
+
33
+ # Define TypedDict for the simplified state
34
+ class HybridGraphState(TypedDict, total=False):
35
+ question: str
36
+ search_kwargs: dict
37
+ docs: list[Document]
38
+ answer: str
39
+
40
+
41
+ def advanced_rag(
42
+ vectorstore: VectorStore | list[VectorStore],
43
+ model_name: str,
44
+ prompt: str | None = None,
45
+ max_tokens: int | None = None,
46
+ temperature: float = 0.1,
47
+ use_preview_models: bool = False,
48
+ enable_query_rewriting: bool = False,
49
+ enable_rrf: bool = False,
50
+ enable_reranking: bool = False,
51
+ # enable_self_evaluation: bool = False,
52
+ query_rewriting_model: str = GROQ_PRODUCTION_MODELS[0],
53
+ query_rewriting_temperature: float = 0.1,
54
+ rerank_model: str = "cohere-rerank-3.5",
55
+ pinecone_api_key: str | None = None,
56
+ rrf_constant: float = 60.0,
57
+ rrf_weights: dict[str, float] | None = None,
58
+ ) -> CompiledStateGraph:
59
+ """
60
+ Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
61
+
62
+ :param prompt: The prompt template to use for the language model. If None, uses chATLAS_Chains.prompt.starters.CHAT_PROMPT_TEMPLATE
63
+ :param vectorstore: Single vectorstore instance or list of vectorstore instances to search
64
+ :param model_name: The name of the language model to use for generating responses.
65
+ :param max_tokens: Maximum number of tokens to generate in the response. If None, uses the model's default value.
66
+ :param temperature: Temperature to use for the model.
67
+ :param use_preview_models: If True, allows the use of preview models from Groq.
68
+ :param enable_query_rewriting: Whether to enable LLM-powered query rewriting.
69
+ :param enable_rrf: Whether to enable RRF (Reciprocal Rank Fusion) for combining results from multiple vectorstores.
70
+ :param enable_reranking: Whether to rerank the retrieved results using the Pinecone API.
71
+ :param query_rewriting_model: Name of the model to use for query rewriting.
72
+ :param query_rewriting_temperature: Temperature to use for query rewriting.
73
+ :param rerank_model: Name of the Pinecone reranker model to use.
74
+ :param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
75
+ :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.
77
+
78
+ :return: A compiled LangGraph with the chosen features
79
+
80
+ """
81
+ if prompt is None:
82
+ prompt = CHAT_PROMPT_TEMPLATE
83
+
84
+ # Create prompt template
85
+ prompt_template = ChatPromptTemplate.from_template(prompt)
86
+
87
+ # Parallel searcher for vectorstore(s)
88
+ searcher = search_runnable(vectorstore)
89
+
90
+ # Initialise models
91
+ model = get_chat_model(model_name, max_tokens, temperature, use_preview_models)
92
+
93
+ if enable_query_rewriting:
94
+ query_rewriting_model_instance = get_chat_model(
95
+ model_name=query_rewriting_model,
96
+ temperature=query_rewriting_temperature,
97
+ use_preview_models=use_preview_models,
98
+ )
99
+
100
+ # --------------- Define functions for the graph nodes ---------------
101
+
102
+ def query_rewriting(state: HybridGraphState) -> HybridGraphState:
103
+ """Rewrite the query to correct typos and enhance clarity."""
104
+ question = state["question"]
105
+
106
+ if not enable_query_rewriting:
107
+ print("[WARNING] Query rewriting disabled, but query_rewriting node was called.")
108
+ # return original state
109
+ return {**state}
110
+
111
+ rewritten_query = rewrite_query(question, model=query_rewriting_model_instance)
112
+
113
+ return {
114
+ **state,
115
+ "question": rewritten_query,
116
+ "unchanged_question": question, # Keep original for reference
117
+ }
118
+
119
+ def retrieval(state: HybridGraphState) -> HybridGraphState:
120
+ """Call the search runnable to retrieve documents"""
121
+
122
+ question = state.get("question")
123
+ if question is None:
124
+ raise Exception("question field is None")
125
+
126
+ search_kwargs = state.get("search_kwargs", {})
127
+ if not search_kwargs:
128
+ print("[WARNING] No search_kwargs provided, using defaults.")
129
+ results = searcher.invoke(question)
130
+ else:
131
+ results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
132
+
133
+ if "docs" not in results:
134
+ raise Exception("Search results missing 'docs' field")
135
+
136
+ print(f"Retrieved {len(results['docs'])} documents")
137
+
138
+ return {
139
+ **state,
140
+ "docs": results["docs"],
141
+ }
142
+
143
+ # Define the document reranking function
144
+ def rerank(state: HybridGraphState) -> HybridGraphState:
145
+ """
146
+ Rerank the parent documents using the Pinecone API.
147
+ """
148
+
149
+ if not enable_reranking:
150
+ print("[WARNING] Reranking disabled, but rerank node was called.")
151
+ # return original state
152
+ return {**state}
153
+
154
+ docs = state.get("docs", [])
155
+
156
+ try:
157
+ reranked_docs = rerank_documents(
158
+ question=state["question"],
159
+ docs=docs,
160
+ reranker_model=rerank_model,
161
+ api_key=pinecone_api_key,
162
+ # num_return_docs = None # return everything
163
+ )
164
+
165
+ except Exception as e:
166
+ print(f"[WARNING] Reranking failed with exception {e}. Returning original documents.")
167
+ reranked_docs = docs
168
+
169
+ return {
170
+ **state,
171
+ "docs": reranked_docs,
172
+ }
173
+
174
+ def rrf(state: HybridGraphState) -> HybridGraphState:
175
+ """Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
176
+ if not enable_rrf:
177
+ print("[WARNING] RRF disabled, but rrf node was called.")
178
+ # return original state
179
+ return {**state}
180
+
181
+ docs = state.get("docs", [])
182
+ if not docs:
183
+ print("[WARNING] No documents retrieved for RRF.")
184
+ return {**state}
185
+
186
+ try:
187
+ rrf_docs = reciprocal_rank_fusion(
188
+ results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
189
+ )
190
+ except Exception as e:
191
+ print(f"[WARNING] RRF failed with exception {e}. Returning original documents.")
192
+ rrf_docs = docs
193
+
194
+ return {
195
+ **state,
196
+ "docs": rrf_docs,
197
+ }
198
+
199
+ def generate_answer(state: HybridGraphState) -> HybridGraphState:
200
+ """Generate a response from the LLM"""
201
+
202
+ # Format the prompt using LangChain template
203
+ prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
204
+ final_prompt = prompt_template.format_messages(**prompt_input)
205
+
206
+ response = model.invoke(final_prompt)
207
+ answer = response.content
208
+
209
+ return {
210
+ **state,
211
+ "answer": answer,
212
+ }
213
+
214
+ # --------------- Build the graph ---------------
215
+ graph = StateGraph(HybridGraphState)
216
+
217
+ # Add all the nodes, but don't link to them if not using
218
+ graph.add_node("query_rewrite", query_rewriting)
219
+ graph.add_node("retrieval", retrieval)
220
+ graph.add_node("rrf", rrf)
221
+ graph.add_node("rerank", rerank)
222
+ graph.add_node("generate", generate_answer)
223
+ # graph.add_node("assess", assess_answer)
224
+ # graph.add_node("refine", refine_answer)
225
+
226
+ if enable_query_rewriting:
227
+ # rewrite the query first
228
+ graph.add_edge("query_rewrite", "retrieval")
229
+ graph.set_entry_point("query_rewrite")
230
+ else:
231
+ # start with retrieval
232
+ graph.set_entry_point("retrieval")
233
+
234
+ if not enable_rrf and not enable_reranking:
235
+ # no document processing, go straight to generation
236
+ graph.add_edge("retrieval", "generate")
237
+
238
+ elif enable_rrf and not enable_reranking:
239
+ # retrieve → rrf → generate
240
+ graph.add_edge("retrieval", "rrf")
241
+ graph.add_edge("rrf", "generate")
242
+
243
+ elif not enable_rrf and enable_reranking:
244
+ # retrieve → rerank → generate
245
+ graph.add_edge("retrieval", "rerank")
246
+ graph.add_edge("rerank", "generate")
247
+
248
+ else:
249
+ # retrieve → rrf → rerank → generate
250
+ graph.add_edge("retrieval", "rrf")
251
+ graph.add_edge("rrf", "rerank")
252
+ graph.add_edge("rerank", "generate")
253
+
254
+ # generate at the end
255
+ graph.add_edge("generate", END)
256
+
257
+ # Compile the graph
258
+ return graph.compile()
259
+
260
+
261
+ if __name__ == "__main__":
262
+ from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
263
+ from chATLAS_Chains.vectorstore import get_vectorstore
264
+
265
+ twiki = get_vectorstore("twiki_prod")
266
+
267
+ # Create the hybrid graph
268
+ graph = advanced_rag(
269
+ vectorstore=[twiki],
270
+ model_name=GROQ_PRODUCTION_MODELS[0],
271
+ enable_query_rewriting=True,
272
+ enable_rrf=True,
273
+ enable_reranking=True,
274
+ )
275
+
276
+ # Test query
277
+ try:
278
+ ans = graph.invoke(
279
+ {
280
+ "question": "How can one check for and remove bad or corrupted events in the analysis?",
281
+ "search_kwargs": {
282
+ "k_text": 3,
283
+ "k": 15,
284
+ "date_filter": "01-01-2010",
285
+ # "type": ["twiki"],
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
297
+
298
+ traceback.print_exc()
299
+ sys.exit(1)
@@ -0,0 +1,82 @@
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.model_selection import get_chat_model
8
+ from chATLAS_Chains.search.basic import search_runnable
9
+ from chATLAS_Chains.utils.doc_utils import combine_documents
10
+ from chATLAS_Embed.Base import VectorStore
11
+
12
+
13
+ def basic_retrieval_chain(
14
+ prompt: str,
15
+ vectorstore: VectorStore | list[VectorStore],
16
+ model_name: str,
17
+ max_tokens: int | None = None,
18
+ temperature: float | None = None,
19
+ ) -> RunnableSerializable:
20
+ """
21
+ Baseline RAG retrieval chain. Searches one or several vectorstores in parallel, passes retrieved documents to the model
22
+
23
+ :param prompt: The prompt template to use with the model.
24
+ :type prompt: str
25
+ :param vectorstore: The vectorstore or list of vectorstores to search over.
26
+ :type vectorstore: Any
27
+ :param model_name: The name of the chat model to use for generating responses.
28
+ :type model_name: str
29
+ :param max_tokens: The maximum number of tokens to generate in the response. Defaults to None.
30
+ :type max_tokens: int | None
31
+ :param temperature: The temperature to use for the model's response generation. Defaults to None.
32
+ :type temperature: float | None
33
+
34
+ :return: A LangChain RunnableSerializable chain that performs retrieval and response generation.
35
+ :rtype: RunnableSerializable
36
+
37
+ """
38
+ prompt_template = ChatPromptTemplate.from_template(prompt)
39
+ model = get_chat_model(model_name, max_tokens, temperature)
40
+
41
+ search = search_runnable(vectorstore)
42
+
43
+ final_inputs = {
44
+ "context": lambda x: combine_documents(x["docs"]),
45
+ "question": itemgetter("question"),
46
+ }
47
+
48
+ answer = {
49
+ "answer": final_inputs | prompt_template | model,
50
+ "docs": lambda x: x["docs"],
51
+ }
52
+
53
+ chain = search | answer
54
+ return chain
55
+
56
+
57
+ if __name__ == "__main__":
58
+ from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
59
+ from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
60
+ from chATLAS_Chains.vectorstore import get_vectorstore
61
+
62
+ twiki_vectorstore = get_vectorstore("twiki_prod")
63
+ mkdocs_vectorstore = get_vectorstore("mkdocs_prod_v1")
64
+
65
+ chain = basic_retrieval_chain(
66
+ prompt=CHAT_PROMPT_TEMPLATE,
67
+ vectorstore=[twiki_vectorstore, mkdocs_vectorstore],
68
+ model_name=GROQ_PRODUCTION_MODELS[0],
69
+ # model_name="meta-llama/llama-4-maverick-17b-128e-instruct",
70
+ # model_name="gemma2-9b-it",
71
+ # model_name="qwen-qwq-32b",
72
+ # model_name="mistral-saba-24b",
73
+ )
74
+ SEARCH_HYPERPARAMS = {"k": 5, "k_text": 0, "date_filter": "01-01-2010"}
75
+
76
+ result = chain.invoke("What is the Higgs boson?", config={"metadata": {"search_kwargs": SEARCH_HYPERPARAMS}})
77
+
78
+ print(f"Answer: {result['answer'].content}")
79
+ print(f"Number of documents retrieved: {len(result['docs'])}")
80
+
81
+ for doc in result["docs"]:
82
+ print(doc.metadata.get("source"))
@@ -4,14 +4,15 @@ Example graph for running langgraph with this general setup and the postgres vec
4
4
 
5
5
  from typing import TypedDict
6
6
 
7
- import langgraph.graph as lg
8
7
  from langchain_core.documents import Document
9
8
  from langchain_core.prompts import ChatPromptTemplate
10
9
  from langgraph.graph import END, StateGraph
10
+ from langgraph.graph.state import CompiledStateGraph
11
11
 
12
+ from chATLAS_Chains.llm.groq import RetryConfig
12
13
  from chATLAS_Chains.llm.model_selection import get_chat_model
13
- from chATLAS_Chains.utils.doc_utils import combine_documents
14
- from chATLAS_Embed import LangChainVectorStore
14
+ from chATLAS_Chains.utils.doc_utils import combine_documents, truncate_to_context_window
15
+ from chATLAS_Embed.LangChainVectorStore import LangChainVectorStore
15
16
 
16
17
 
17
18
  # Define TypedDict for the state
@@ -24,7 +25,15 @@ class GraphState(TypedDict, total=False):
24
25
  answer: str
25
26
 
26
27
 
27
- def basic_retrieval_graph(prompt: str, vectorstore, model_name: str) -> lg.Graph:
28
+ def basic_retrieval_graph(
29
+ prompt: str,
30
+ vectorstore,
31
+ model_name: str,
32
+ max_tokens: int | None = None,
33
+ temperature: float | None = None,
34
+ use_preview_models: bool = False,
35
+ retry_config: RetryConfig | None = None,
36
+ ) -> CompiledStateGraph:
28
37
  """
29
38
  Baseline RAG retrieval graph using LangGraph. Searches one or several vectorstores,
30
39
  passes retrieved documents to the model, and returns the final answer.
@@ -38,7 +47,7 @@ def basic_retrieval_graph(prompt: str, vectorstore, model_name: str) -> lg.Graph
38
47
  A LangGraph graph that can be executed for RAG
39
48
  """
40
49
  # Initialize the model and prompt template
41
- model = get_chat_model(model_name)
50
+ model = get_chat_model(model_name, use_preview_models=use_preview_models)
42
51
  prompt_template = ChatPromptTemplate.from_template(prompt)
43
52
 
44
53
  # Create a list of retrievers from the vectorstore(s)
@@ -84,8 +93,16 @@ def basic_retrieval_graph(prompt: str, vectorstore, model_name: str) -> lg.Graph
84
93
  # Define the document processing function
85
94
  def process_docs(state: GraphState) -> GraphState:
86
95
  """Process the merged documents into a context string."""
87
- docs = state["merged_docs"]
88
- context = combine_documents(docs)
96
+ initial_docs = state["merged_docs"]
97
+
98
+ truncated_docs = truncate_to_context_window(
99
+ docs=initial_docs, question=state["question"], prompt=prompt_template, model_name=model_name
100
+ )
101
+
102
+ if len(truncated_docs) < len(initial_docs):
103
+ print(f"Truncated from {len(initial_docs)} to {len(truncated_docs)} docs to fit context window.")
104
+
105
+ context = combine_documents(truncated_docs)
89
106
  return {
90
107
  "question": state["question"],
91
108
  "retrieved_docs": state["retrieved_docs"],
@@ -136,23 +153,26 @@ def basic_retrieval_graph(prompt: str, vectorstore, model_name: str) -> lg.Graph
136
153
 
137
154
 
138
155
  if __name__ == "__main__":
139
- # Example of how to run the graph correctly
140
- import os
156
+ # Example usage
157
+ from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
158
+ from chATLAS_Chains.vectorstore import get_vectorstore
141
159
 
142
- os.environ["CHATLAS_EMBEDDING_MODEL_PATH"] = "<PATH TO YOUR EMBEDDING MODEL>"
143
- os.environ["CHATLAS_OPENAI_KEY"] = "YOUR OPENAI API KEY"
144
- os.environ["CHATLAS_DB_PASSWORD"] = "<>"
160
+ vs = get_vectorstore("twiki_prod")
145
161
 
146
- from ..prompt.starters import CHAT_PROMPT_TEMPLATE
147
- from ..vectorstore import vectorstore
162
+ # model_name = "llama-3.3-70b-versatile"
163
+ model_name = "deepseek-r1-distill-llama-70b"
164
+ # model_name="gemma2-9b-it" # smaller context window, to check the doc truncation
165
+ # model_name = "gpt-4o-mini"
148
166
 
149
- graph = basic_retrieval_graph(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vectorstore, model_name="gpt-4o-mini")
167
+ graph = basic_retrieval_graph(
168
+ prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vs, model_name=model_name, use_preview_models=True
169
+ )
150
170
 
151
- ans = graph.invoke(
171
+ result = graph.invoke(
152
172
  {
153
- "question": "How many onions are in ATLAS",
173
+ "question": "What is the crack veto for electron reconstruction",
154
174
  "search_kwargs": {
155
- "k_text": 3,
175
+ "k_text": 5,
156
176
  "k": 10,
157
177
  "date_filter": "01-01-2010",
158
178
  "type": ["CDS", "twiki", "Indico"],
@@ -160,4 +180,5 @@ if __name__ == "__main__":
160
180
  }
161
181
  )
162
182
 
163
- print(ans)
183
+ print("Num retrieved documents: ", len(result["merged_docs"]))
184
+ print("Answer: ", result["answer"])