chATLAS_Chains 0.1.6__tar.gz → 0.2.0__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.6 → chatlas_chains-0.2.0}/PKG-INFO +72 -7
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/README.md +71 -6
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/advanced.py +66 -54
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/basic.py +9 -18
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/basic_graph.py +2 -6
- chatlas_chains-0.2.0/chATLAS_Chains/chains/conversational_graph.py +575 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/documents/rerank.py +2 -1
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/llm/groq.py +113 -4
- chatlas_chains-0.2.0/chATLAS_Chains/llm/model_selection.py +302 -0
- chatlas_chains-0.2.0/chATLAS_Chains/llm/runnables.py +156 -0
- chatlas_chains-0.2.0/chATLAS_Chains/prompt/starters.py +314 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/search/basic.py +1 -1
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/vectorstore.py +4 -3
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/PKG-INFO +72 -7
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/SOURCES.txt +6 -1
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/requires.txt +1 -1
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/pyproject.toml +2 -2
- chatlas_chains-0.2.0/tests/conftest.py +99 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_chains.py +11 -15
- chatlas_chains-0.2.0/tests/test_chat_model_kwargs.py +170 -0
- chatlas_chains-0.2.0/tests/test_conversational.py +238 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_groq.py +23 -5
- chatlas_chains-0.2.0/tests/test_llm_runnables.py +48 -0
- chatlas_chains-0.2.0/tests/test_model_selection.py +129 -0
- chatlas_chains-0.1.6/chATLAS_Chains/llm/model_selection.py +0 -176
- chatlas_chains-0.1.6/chATLAS_Chains/prompt/starters.py +0 -149
- chatlas_chains-0.1.6/tests/conftest.py +0 -68
- chatlas_chains-0.1.6/tests/test_llm.py +0 -227
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/LICENSE +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/__init__.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/__init__.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/documents/rrf.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/llm/__init__.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/log.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/prompt/__init__.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/query/query_rewriting.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/search/__init__.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/utils/__init__.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/utils/doc_utils.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/top_level.txt +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/setup.cfg +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/__init__.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_rrf.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_search.py +0 -0
- {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: chATLAS_Chains
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.2.0
|
|
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
|
|
@@ -19,9 +19,9 @@ Requires-Dist: langgraph
|
|
|
19
19
|
Requires-Dist: sentence-transformers>=3.0.0
|
|
20
20
|
Requires-Dist: tiktoken
|
|
21
21
|
Requires-Dist: pinecone
|
|
22
|
-
Requires-Dist: torch==2.2.1
|
|
23
22
|
Requires-Dist: chatlas-embed
|
|
24
23
|
Requires-Dist: psycopg2-binary>=2.9.10
|
|
24
|
+
Requires-Dist: httpx[socks]>=0.28.1
|
|
25
25
|
Dynamic: license-file
|
|
26
26
|
|
|
27
27
|
|
|
@@ -62,14 +62,71 @@ export CHATLAS_GROQ_KEY="your groq api key"
|
|
|
62
62
|
|
|
63
63
|
**note** The API address is local to the CERN network. If not at CERN, you can forward it like so:
|
|
64
64
|
```bash
|
|
65
|
-
ssh -L 3000:cs-513-ml003:3000
|
|
65
|
+
ssh -L 3000:cs-513-ml003:3000 $LXPLUS_USERNAME@lxplus.cern.ch
|
|
66
66
|
export CHATLAS_GROQ_BASE_URL="http://localhost:3000"
|
|
67
67
|
```
|
|
68
68
|
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
|
|
69
|
+
3. Using LLMs via CERN's LiteLLM API, here is the [repo](https://gitlab.cern.ch/itgpt/litellm-okd/-/tree/main) and some [setup instructions](https://codimd.web.cern.ch/tQKiMa13Q4O-EJXWTO3N7w?view#Using-Your-Dedicated-API-Key-to-Access-LLMs) for reference.
|
|
70
|
+
```bash
|
|
71
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your litellm key"
|
|
72
|
+
```
|
|
73
|
+
|
|
74
|
+
## Supported Chains
|
|
75
|
+
|
|
76
|
+
More details [here](chATLAS_Chains/chains/README.md)
|
|
77
|
+
|
|
78
|
+
- `chains.basic.basic_retrieval_chain`
|
|
79
|
+
- `chains.advanced.advanced_rag`
|
|
80
|
+
|
|
81
|
+
### Model Configuration in Chains
|
|
82
|
+
|
|
83
|
+
Supported chain constructors now accept a typed `chat_model_kwargs` argument for model options (for example:
|
|
84
|
+
`temperature`, `max_tokens`, `service_provider`, `api_key`, `base_url`, `proxy`).
|
|
85
|
+
|
|
86
|
+
```python
|
|
87
|
+
from chATLAS_Chains.chains.basic import basic_retrieval_chain
|
|
88
|
+
|
|
89
|
+
chain = basic_retrieval_chain(
|
|
90
|
+
prompt=...,
|
|
91
|
+
vectorstore=...,
|
|
92
|
+
model_name="gpt-4o-mini",
|
|
93
|
+
chat_model_kwargs={"temperature": 0.1, "max_tokens": 512},
|
|
94
|
+
)
|
|
95
|
+
```
|
|
96
|
+
|
|
97
|
+
## Forwarding vectorstore connections
|
|
98
|
+
|
|
99
|
+
If not on the CERN network, you can forward the connection to the postgres servers with:
|
|
100
|
+
|
|
101
|
+
```bash
|
|
102
|
+
ssh -N \
|
|
103
|
+
-L 6624:dbod-chatlas.cern.ch:6624 \
|
|
104
|
+
-L 6606:dbod-chatlas-cds.cern.ch:6606 \
|
|
105
|
+
"$LXPLUS_USERNAME"@lxplus.cern.ch
|
|
106
|
+
export CHATLAS_PORT_FORWARDING=1
|
|
107
|
+
```
|
|
108
|
+
|
|
109
|
+
You can then the helper function [`get_vectorstore`](chATLAS_Chains/vectorstore.py)
|
|
110
|
+
|
|
111
|
+
## Testing Environment Variables
|
|
112
|
+
|
|
113
|
+
Some tests are DB-backed integration tests (`tests/test_chains.py`, `tests/test_conversational.py`, `tests/test_search.py`).
|
|
114
|
+
If the DB/test environment is not configured, these tests are skipped by `tests/conftest.py`.
|
|
115
|
+
|
|
116
|
+
`tests/conftest.py` now uses explicit controls:
|
|
117
|
+
|
|
118
|
+
- `CHATLAS_PORT_FORWARDING`: enable localhost DB tunnels (`1`, `true`, `True`)
|
|
119
|
+
- `CHATLAS_DB_PASSWORD`
|
|
120
|
+
|
|
121
|
+
### Local Example (with DB tunnels)
|
|
122
|
+
|
|
123
|
+
```bash
|
|
124
|
+
export CHATLAS_DB_PASSWORD="..."
|
|
125
|
+
export CHATLAS_PORT_FORWARDING=1
|
|
126
|
+
unset GITLAB_PAT
|
|
127
|
+
|
|
128
|
+
uv run pytest -q
|
|
129
|
+
```
|
|
73
130
|
|
|
74
131
|
## Postgres
|
|
75
132
|
|
|
@@ -91,6 +148,14 @@ CREATE EXTENSION IF NOT EXISTS vector;
|
|
|
91
148
|
```
|
|
92
149
|
## CHANGELOG
|
|
93
150
|
|
|
151
|
+
#### 0.1.7
|
|
152
|
+
|
|
153
|
+
Support for CERN-hosted LiteLLM models
|
|
154
|
+
|
|
155
|
+
Multi-turn conversational RAG with (local) conversation history
|
|
156
|
+
|
|
157
|
+
Bugfixes
|
|
158
|
+
|
|
94
159
|
#### 0.1.6
|
|
95
160
|
|
|
96
161
|
Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
|
|
@@ -36,14 +36,71 @@ export CHATLAS_GROQ_KEY="your groq api key"
|
|
|
36
36
|
|
|
37
37
|
**note** The API address is local to the CERN network. If not at CERN, you can forward it like so:
|
|
38
38
|
```bash
|
|
39
|
-
ssh -L 3000:cs-513-ml003:3000
|
|
39
|
+
ssh -L 3000:cs-513-ml003:3000 $LXPLUS_USERNAME@lxplus.cern.ch
|
|
40
40
|
export CHATLAS_GROQ_BASE_URL="http://localhost:3000"
|
|
41
41
|
```
|
|
42
42
|
|
|
43
|
-
|
|
44
|
-
|
|
45
|
-
|
|
46
|
-
|
|
43
|
+
3. Using LLMs via CERN's LiteLLM API, here is the [repo](https://gitlab.cern.ch/itgpt/litellm-okd/-/tree/main) and some [setup instructions](https://codimd.web.cern.ch/tQKiMa13Q4O-EJXWTO3N7w?view#Using-Your-Dedicated-API-Key-to-Access-LLMs) for reference.
|
|
44
|
+
```bash
|
|
45
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your litellm key"
|
|
46
|
+
```
|
|
47
|
+
|
|
48
|
+
## Supported Chains
|
|
49
|
+
|
|
50
|
+
More details [here](chATLAS_Chains/chains/README.md)
|
|
51
|
+
|
|
52
|
+
- `chains.basic.basic_retrieval_chain`
|
|
53
|
+
- `chains.advanced.advanced_rag`
|
|
54
|
+
|
|
55
|
+
### Model Configuration in Chains
|
|
56
|
+
|
|
57
|
+
Supported chain constructors now accept a typed `chat_model_kwargs` argument for model options (for example:
|
|
58
|
+
`temperature`, `max_tokens`, `service_provider`, `api_key`, `base_url`, `proxy`).
|
|
59
|
+
|
|
60
|
+
```python
|
|
61
|
+
from chATLAS_Chains.chains.basic import basic_retrieval_chain
|
|
62
|
+
|
|
63
|
+
chain = basic_retrieval_chain(
|
|
64
|
+
prompt=...,
|
|
65
|
+
vectorstore=...,
|
|
66
|
+
model_name="gpt-4o-mini",
|
|
67
|
+
chat_model_kwargs={"temperature": 0.1, "max_tokens": 512},
|
|
68
|
+
)
|
|
69
|
+
```
|
|
70
|
+
|
|
71
|
+
## Forwarding vectorstore connections
|
|
72
|
+
|
|
73
|
+
If not on the CERN network, you can forward the connection to the postgres servers with:
|
|
74
|
+
|
|
75
|
+
```bash
|
|
76
|
+
ssh -N \
|
|
77
|
+
-L 6624:dbod-chatlas.cern.ch:6624 \
|
|
78
|
+
-L 6606:dbod-chatlas-cds.cern.ch:6606 \
|
|
79
|
+
"$LXPLUS_USERNAME"@lxplus.cern.ch
|
|
80
|
+
export CHATLAS_PORT_FORWARDING=1
|
|
81
|
+
```
|
|
82
|
+
|
|
83
|
+
You can then the helper function [`get_vectorstore`](chATLAS_Chains/vectorstore.py)
|
|
84
|
+
|
|
85
|
+
## Testing Environment Variables
|
|
86
|
+
|
|
87
|
+
Some tests are DB-backed integration tests (`tests/test_chains.py`, `tests/test_conversational.py`, `tests/test_search.py`).
|
|
88
|
+
If the DB/test environment is not configured, these tests are skipped by `tests/conftest.py`.
|
|
89
|
+
|
|
90
|
+
`tests/conftest.py` now uses explicit controls:
|
|
91
|
+
|
|
92
|
+
- `CHATLAS_PORT_FORWARDING`: enable localhost DB tunnels (`1`, `true`, `True`)
|
|
93
|
+
- `CHATLAS_DB_PASSWORD`
|
|
94
|
+
|
|
95
|
+
### Local Example (with DB tunnels)
|
|
96
|
+
|
|
97
|
+
```bash
|
|
98
|
+
export CHATLAS_DB_PASSWORD="..."
|
|
99
|
+
export CHATLAS_PORT_FORWARDING=1
|
|
100
|
+
unset GITLAB_PAT
|
|
101
|
+
|
|
102
|
+
uv run pytest -q
|
|
103
|
+
```
|
|
47
104
|
|
|
48
105
|
## Postgres
|
|
49
106
|
|
|
@@ -65,6 +122,14 @@ CREATE EXTENSION IF NOT EXISTS vector;
|
|
|
65
122
|
```
|
|
66
123
|
## CHANGELOG
|
|
67
124
|
|
|
125
|
+
#### 0.1.7
|
|
126
|
+
|
|
127
|
+
Support for CERN-hosted LiteLLM models
|
|
128
|
+
|
|
129
|
+
Multi-turn conversational RAG with (local) conversation history
|
|
130
|
+
|
|
131
|
+
Bugfixes
|
|
132
|
+
|
|
68
133
|
#### 0.1.6
|
|
69
134
|
|
|
70
135
|
Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
|
|
@@ -119,4 +184,4 @@ chATLAS_Benchmark is released under Apache v2.0 license.
|
|
|
119
184
|
|
|
120
185
|
*For questions and support, please [contact](mailto:joseph.caimin.egan@cern.ch)*
|
|
121
186
|
|
|
122
|
-
</div>
|
|
187
|
+
</div>
|
|
@@ -23,13 +23,17 @@ from langgraph.graph.state import CompiledStateGraph
|
|
|
23
23
|
|
|
24
24
|
from chATLAS_Chains.documents.rerank import rerank_documents
|
|
25
25
|
from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
|
|
26
|
-
from chATLAS_Chains.llm.
|
|
27
|
-
|
|
26
|
+
from chATLAS_Chains.llm.model_selection import (
|
|
27
|
+
GROQ_PRODUCTION_MODELS,
|
|
28
|
+
ChatModelKwargs,
|
|
29
|
+
get_chat_model,
|
|
30
|
+
sanitize_chat_model_kwargs,
|
|
31
|
+
)
|
|
28
32
|
from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
|
|
29
33
|
from chATLAS_Chains.query.query_rewriting import rewrite_query
|
|
30
34
|
from chATLAS_Chains.search.basic import search_runnable
|
|
31
35
|
from chATLAS_Chains.utils.doc_utils import combine_documents
|
|
32
|
-
from chATLAS_Embed.
|
|
36
|
+
from chATLAS_Embed.VectorStores import VectorStore
|
|
33
37
|
|
|
34
38
|
logger = logging.getLogger(__name__)
|
|
35
39
|
|
|
@@ -40,27 +44,25 @@ class HybridGraphState(TypedDict, total=False):
|
|
|
40
44
|
search_kwargs: dict
|
|
41
45
|
docs: list[Document]
|
|
42
46
|
answer: str
|
|
47
|
+
chat_history: str # Optional
|
|
43
48
|
|
|
44
49
|
|
|
45
50
|
def advanced_rag(
|
|
46
51
|
vectorstore: VectorStore | list[VectorStore],
|
|
47
52
|
model_name: str,
|
|
48
53
|
prompt: str | None = None,
|
|
49
|
-
|
|
50
|
-
temperature: float = 0.1,
|
|
51
|
-
use_preview_models: bool = False,
|
|
54
|
+
chat_model_kwargs: ChatModelKwargs | None = None,
|
|
52
55
|
enable_query_rewriting: bool = False,
|
|
53
56
|
enable_rrf: bool = False,
|
|
54
57
|
enable_reranking: bool = False,
|
|
55
58
|
# enable_self_evaluation: bool = False,
|
|
56
59
|
query_rewriting_model: str = GROQ_PRODUCTION_MODELS[0],
|
|
57
|
-
|
|
60
|
+
query_rewriting_chat_model_kwargs: ChatModelKwargs | None = None,
|
|
58
61
|
rerank_model: str = "cohere-rerank-3.5",
|
|
59
62
|
pinecone_api_key: str | None = None,
|
|
60
63
|
rrf_constant: float = 60.0,
|
|
61
64
|
rrf_weights: dict[str, float] | None = None,
|
|
62
|
-
|
|
63
|
-
fallback_models: list[str] | None = None,
|
|
65
|
+
skip_generation: bool = False,
|
|
64
66
|
) -> CompiledStateGraph:
|
|
65
67
|
"""
|
|
66
68
|
Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
|
|
@@ -68,20 +70,17 @@ def advanced_rag(
|
|
|
68
70
|
:param prompt: The prompt template to use for the language model. If None, uses chATLAS_Chains.prompt.starters.CHAT_PROMPT_TEMPLATE
|
|
69
71
|
:param vectorstore: Single vectorstore instance or list of vectorstore instances to search
|
|
70
72
|
:param model_name: The name of the language model to use for generating responses.
|
|
71
|
-
:param
|
|
72
|
-
:param temperature: Temperature to use for the model.
|
|
73
|
-
:param use_preview_models: If True, allows the use of preview models from Groq.
|
|
73
|
+
:param chat_model_kwargs: Optional kwargs passed through to ``get_chat_model`` for the main model.
|
|
74
74
|
:param enable_query_rewriting: Whether to enable LLM-powered query rewriting.
|
|
75
75
|
:param enable_rrf: Whether to enable RRF (Reciprocal Rank Fusion) for combining results from multiple vectorstores.
|
|
76
76
|
:param enable_reranking: Whether to rerank the retrieved results using the Pinecone API.
|
|
77
77
|
:param query_rewriting_model: Name of the model to use for query rewriting.
|
|
78
|
-
:param
|
|
78
|
+
:param query_rewriting_chat_model_kwargs: Optional kwargs for the query rewriting model.
|
|
79
|
+
Inheritance order is: ``chat_model_kwargs`` -> ``{"temperature": 0.1}`` defaults -> explicit query rewriting kwargs.
|
|
79
80
|
:param rerank_model: Name of the Pinecone reranker model to use.
|
|
80
81
|
:param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
|
|
81
82
|
:param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
|
|
82
83
|
: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.
|
|
85
84
|
|
|
86
85
|
:return: A compiled LangGraph with the chosen features
|
|
87
86
|
|
|
@@ -95,17 +94,18 @@ def advanced_rag(
|
|
|
95
94
|
# Parallel searcher for vectorstore(s)
|
|
96
95
|
searcher = search_runnable(vectorstore)
|
|
97
96
|
|
|
97
|
+
primary_model_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
|
|
98
|
+
|
|
98
99
|
# Initialise models
|
|
99
|
-
model = get_chat_model(model_name,
|
|
100
|
+
model = get_chat_model(model_name, **primary_model_kwargs)
|
|
100
101
|
|
|
101
102
|
if enable_query_rewriting:
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
temperature
|
|
105
|
-
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
)
|
|
103
|
+
query_model_kwargs: ChatModelKwargs = {
|
|
104
|
+
**primary_model_kwargs,
|
|
105
|
+
"temperature": 0.1,
|
|
106
|
+
}
|
|
107
|
+
query_model_kwargs.update(sanitize_chat_model_kwargs(query_rewriting_chat_model_kwargs))
|
|
108
|
+
query_rewriting_model_instance = get_chat_model(model_name=query_rewriting_model, **query_model_kwargs)
|
|
109
109
|
|
|
110
110
|
# --------------- Define functions for the graph nodes ---------------
|
|
111
111
|
|
|
@@ -209,7 +209,12 @@ def advanced_rag(
|
|
|
209
209
|
def generate_answer(state: HybridGraphState) -> HybridGraphState:
|
|
210
210
|
"""Generate a response from the LLM"""
|
|
211
211
|
# Format the prompt using LangChain template
|
|
212
|
-
prompt_input = {
|
|
212
|
+
prompt_input = {
|
|
213
|
+
"context": combine_documents(state["docs"]),
|
|
214
|
+
"question": state["question"],
|
|
215
|
+
"chat_history": state.get("chat_history", ""),
|
|
216
|
+
}
|
|
217
|
+
|
|
213
218
|
final_prompt = prompt_template.format_messages(**prompt_input)
|
|
214
219
|
|
|
215
220
|
response = model.invoke(final_prompt)
|
|
@@ -223,50 +228,62 @@ def advanced_rag(
|
|
|
223
228
|
# --------------- Build the graph ---------------
|
|
224
229
|
graph = StateGraph(HybridGraphState)
|
|
225
230
|
|
|
226
|
-
# Add all the nodes, but don't link to them if not using
|
|
227
231
|
graph.add_node("query_rewrite", query_rewriting)
|
|
228
232
|
graph.add_node("retrieval", retrieval)
|
|
229
233
|
graph.add_node("rrf", rrf)
|
|
230
234
|
graph.add_node("rerank", rerank)
|
|
231
|
-
|
|
232
|
-
|
|
233
|
-
# graph.add_node("refine", refine_answer)
|
|
235
|
+
if not skip_generation:
|
|
236
|
+
graph.add_node("generate", generate_answer)
|
|
234
237
|
|
|
235
238
|
if enable_query_rewriting:
|
|
236
|
-
# rewrite the query first
|
|
237
239
|
graph.add_edge("query_rewrite", "retrieval")
|
|
238
240
|
graph.set_entry_point("query_rewrite")
|
|
239
241
|
else:
|
|
240
|
-
# start with retrieval
|
|
241
242
|
graph.set_entry_point("retrieval")
|
|
242
243
|
|
|
243
|
-
|
|
244
|
-
|
|
245
|
-
|
|
246
|
-
|
|
247
|
-
elif enable_rrf and not enable_reranking:
|
|
248
|
-
# retrieve → rrf → generate
|
|
244
|
+
# Determine the last pre-generation node
|
|
245
|
+
if enable_rrf and enable_reranking:
|
|
246
|
+
last_node = "rerank"
|
|
249
247
|
graph.add_edge("retrieval", "rrf")
|
|
250
|
-
graph.add_edge("rrf", "
|
|
251
|
-
|
|
252
|
-
|
|
253
|
-
|
|
248
|
+
graph.add_edge("rrf", "rerank")
|
|
249
|
+
elif enable_rrf:
|
|
250
|
+
last_node = "rrf"
|
|
251
|
+
graph.add_edge("retrieval", "rrf")
|
|
252
|
+
elif enable_reranking:
|
|
253
|
+
last_node = "rerank"
|
|
254
254
|
graph.add_edge("retrieval", "rerank")
|
|
255
|
-
graph.add_edge("rerank", "generate")
|
|
256
|
-
|
|
257
255
|
else:
|
|
258
|
-
|
|
259
|
-
graph.add_edge("retrieval", "rrf")
|
|
260
|
-
graph.add_edge("rrf", "rerank")
|
|
261
|
-
graph.add_edge("rerank", "generate")
|
|
256
|
+
last_node = "retrieval"
|
|
262
257
|
|
|
263
|
-
|
|
264
|
-
|
|
258
|
+
if skip_generation:
|
|
259
|
+
graph.add_edge(last_node, END)
|
|
260
|
+
else:
|
|
261
|
+
graph.add_edge(last_node, "generate")
|
|
262
|
+
graph.add_edge("generate", END)
|
|
265
263
|
|
|
266
|
-
# Compile the graph
|
|
267
264
|
return graph.compile()
|
|
268
265
|
|
|
269
266
|
|
|
267
|
+
def build_generation_prompt(prompt: str | None = None) -> ChatPromptTemplate:
|
|
268
|
+
"""Build the prompt template used by the generation step."""
|
|
269
|
+
if prompt is None:
|
|
270
|
+
prompt = CHAT_PROMPT_TEMPLATE
|
|
271
|
+
return ChatPromptTemplate.from_template(prompt)
|
|
272
|
+
|
|
273
|
+
|
|
274
|
+
def stream_generate_answer(state: HybridGraphState, model, prompt_template: ChatPromptTemplate):
|
|
275
|
+
"""Generator that streams LLM tokens for the generation step."""
|
|
276
|
+
prompt_input = {
|
|
277
|
+
"context": combine_documents(state["docs"]),
|
|
278
|
+
"question": state["question"],
|
|
279
|
+
"chat_history": state.get("chat_history", ""),
|
|
280
|
+
}
|
|
281
|
+
final_prompt = prompt_template.format_messages(**prompt_input)
|
|
282
|
+
for chunk in model.stream(final_prompt):
|
|
283
|
+
if chunk.content:
|
|
284
|
+
yield chunk.content
|
|
285
|
+
|
|
286
|
+
|
|
270
287
|
if __name__ == "__main__":
|
|
271
288
|
from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
|
|
272
289
|
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
@@ -274,11 +291,6 @@ if __name__ == "__main__":
|
|
|
274
291
|
twiki = get_vectorstore("twiki_prod")
|
|
275
292
|
mkdocs = get_vectorstore("mkdocs_prod_v1")
|
|
276
293
|
|
|
277
|
-
retry_config = RetryConfig(
|
|
278
|
-
max_retries=1,
|
|
279
|
-
max_delay=120.0,
|
|
280
|
-
)
|
|
281
|
-
|
|
282
294
|
# Create the hybrid graph
|
|
283
295
|
graph = advanced_rag(
|
|
284
296
|
vectorstore=[twiki, mkdocs],
|
|
@@ -1,24 +1,19 @@
|
|
|
1
1
|
from operator import itemgetter
|
|
2
|
-
from typing import Optional
|
|
3
2
|
|
|
4
3
|
from langchain_core.prompts import ChatPromptTemplate
|
|
5
4
|
from langchain_core.runnables import RunnableSerializable
|
|
6
5
|
|
|
7
|
-
from chATLAS_Chains.llm.
|
|
8
|
-
from chATLAS_Chains.llm.model_selection import get_chat_model
|
|
6
|
+
from chATLAS_Chains.llm.model_selection import ChatModelKwargs, get_chat_model, sanitize_chat_model_kwargs
|
|
9
7
|
from chATLAS_Chains.search.basic import search_runnable
|
|
10
8
|
from chATLAS_Chains.utils.doc_utils import combine_documents
|
|
11
|
-
from chATLAS_Embed.
|
|
9
|
+
from chATLAS_Embed.VectorStores import VectorStore
|
|
12
10
|
|
|
13
11
|
|
|
14
12
|
def basic_retrieval_chain(
|
|
15
13
|
prompt: str,
|
|
16
14
|
vectorstore: VectorStore | list[VectorStore],
|
|
17
15
|
model_name: str,
|
|
18
|
-
|
|
19
|
-
temperature: float | None = None,
|
|
20
|
-
use_preview_models: bool = False,
|
|
21
|
-
retry_config: RetryConfig | None = None,
|
|
16
|
+
chat_model_kwargs: ChatModelKwargs | None = None,
|
|
22
17
|
) -> RunnableSerializable:
|
|
23
18
|
"""
|
|
24
19
|
Baseline RAG retrieval chain. Searches one or several vectorstores in parallel, passes retrieved documents to the model
|
|
@@ -29,21 +24,16 @@ def basic_retrieval_chain(
|
|
|
29
24
|
:type vectorstore: Any
|
|
30
25
|
:param model_name: The name of the chat model to use for generating responses.
|
|
31
26
|
:type model_name: str
|
|
32
|
-
:param
|
|
33
|
-
:type
|
|
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
|
|
27
|
+
:param chat_model_kwargs: Optional kwargs passed through to ``get_chat_model``.
|
|
28
|
+
:type chat_model_kwargs: ChatModelKwargs | None
|
|
40
29
|
|
|
41
30
|
:return: A LangChain RunnableSerializable chain that performs retrieval and response generation.
|
|
42
31
|
:rtype: RunnableSerializable
|
|
43
32
|
|
|
44
33
|
"""
|
|
45
34
|
prompt_template = ChatPromptTemplate.from_template(prompt)
|
|
46
|
-
|
|
35
|
+
sanitized_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
|
|
36
|
+
model = get_chat_model(model_name, **sanitized_kwargs)
|
|
47
37
|
|
|
48
38
|
search = search_runnable(vectorstore)
|
|
49
39
|
|
|
@@ -72,7 +62,8 @@ if __name__ == "__main__":
|
|
|
72
62
|
chain = basic_retrieval_chain(
|
|
73
63
|
prompt=CHAT_PROMPT_TEMPLATE,
|
|
74
64
|
vectorstore=[twiki_vectorstore, mkdocs_vectorstore],
|
|
75
|
-
model_name=
|
|
65
|
+
model_name="openai/gpt-oss-120b",
|
|
66
|
+
# chat_model_kwargs={"service_provider": "groq"},
|
|
76
67
|
# model_name="meta-llama/llama-4-maverick-17b-128e-instruct",
|
|
77
68
|
# model_name="gemma2-9b-it",
|
|
78
69
|
# model_name="qwen-qwq-32b",
|
|
@@ -31,8 +31,6 @@ def basic_retrieval_graph(
|
|
|
31
31
|
model_name: str,
|
|
32
32
|
max_tokens: int | None = None,
|
|
33
33
|
temperature: float | None = None,
|
|
34
|
-
use_preview_models: bool = False,
|
|
35
|
-
retry_config: RetryConfig | None = None,
|
|
36
34
|
) -> CompiledStateGraph:
|
|
37
35
|
"""
|
|
38
36
|
Baseline RAG retrieval graph using LangGraph. Searches one or several vectorstores,
|
|
@@ -47,7 +45,7 @@ def basic_retrieval_graph(
|
|
|
47
45
|
A LangGraph graph that can be executed for RAG
|
|
48
46
|
"""
|
|
49
47
|
# Initialize the model and prompt template
|
|
50
|
-
model = get_chat_model(model_name
|
|
48
|
+
model = get_chat_model(model_name)
|
|
51
49
|
prompt_template = ChatPromptTemplate.from_template(prompt)
|
|
52
50
|
|
|
53
51
|
# Create a list of retrievers from the vectorstore(s)
|
|
@@ -164,9 +162,7 @@ if __name__ == "__main__":
|
|
|
164
162
|
# model_name="gemma2-9b-it" # smaller context window, to check the doc truncation
|
|
165
163
|
# model_name = "gpt-4o-mini"
|
|
166
164
|
|
|
167
|
-
graph = basic_retrieval_graph(
|
|
168
|
-
prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vs, model_name=model_name, use_preview_models=True
|
|
169
|
-
)
|
|
165
|
+
graph = basic_retrieval_graph(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vs, model_name=model_name)
|
|
170
166
|
|
|
171
167
|
result = graph.invoke(
|
|
172
168
|
{
|