chATLAS_Chains 0.1.5__tar.gz → 0.1.7__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/PKG-INFO +78 -7
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/README.md +77 -6
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/advanced.py +56 -60
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/basic.py +9 -18
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/basic_graph.py +2 -6
- chatlas_chains-0.1.7/chATLAS_Chains/chains/conversational_graph.py +383 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/documents/rerank.py +2 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/documents/rrf.py +29 -26
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/llm/groq.py +3 -2
- chatlas_chains-0.1.7/chATLAS_Chains/llm/model_selection.py +302 -0
- chatlas_chains-0.1.7/chATLAS_Chains/llm/runnables.py +156 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/starters.py +136 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/search/basic.py +1 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/vectorstore.py +4 -3
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/PKG-INFO +78 -7
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/SOURCES.txt +7 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/requires.txt +1 -1
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/pyproject.toml +2 -2
- chatlas_chains-0.1.7/tests/conftest.py +99 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/test_chains.py +41 -33
- chatlas_chains-0.1.7/tests/test_chat_model_kwargs.py +168 -0
- chatlas_chains-0.1.7/tests/test_conversational.py +238 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/test_groq.py +0 -1
- chatlas_chains-0.1.7/tests/test_llm_runnables.py +48 -0
- chatlas_chains-0.1.7/tests/test_model_selection.py +129 -0
- chatlas_chains-0.1.7/tests/test_rrf.py +126 -0
- chatlas_chains-0.1.5/chATLAS_Chains/llm/model_selection.py +0 -172
- chatlas_chains-0.1.5/tests/conftest.py +0 -68
- chatlas_chains-0.1.5/tests/test_llm.py +0 -205
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/LICENSE +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/llm/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/log.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/query/query_rewriting.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/search/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/utils/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/utils/doc_utils.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/top_level.txt +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/setup.cfg +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/__init__.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/test_search.py +0 -0
- {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/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.7
|
|
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,20 @@ 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
|
+
|
|
159
|
+
#### 0.1.6
|
|
160
|
+
|
|
161
|
+
Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
|
|
162
|
+
|
|
163
|
+
Add `fallback_models` optional argument to `advanced_rag`
|
|
164
|
+
|
|
94
165
|
#### 0.1.5
|
|
95
166
|
|
|
96
167
|
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
@@ -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,20 @@ 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
|
+
|
|
133
|
+
#### 0.1.6
|
|
134
|
+
|
|
135
|
+
Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
|
|
136
|
+
|
|
137
|
+
Add `fallback_models` optional argument to `advanced_rag`
|
|
138
|
+
|
|
68
139
|
#### 0.1.5
|
|
69
140
|
|
|
70
141
|
Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
|
|
@@ -113,4 +184,4 @@ chATLAS_Benchmark is released under Apache v2.0 license.
|
|
|
113
184
|
|
|
114
185
|
*For questions and support, please [contact](mailto:joseph.caimin.egan@cern.ch)*
|
|
115
186
|
|
|
116
|
-
</div>
|
|
187
|
+
</div>
|
|
@@ -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,13 +23,19 @@ 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
|
|
25
|
-
from chATLAS_Chains.llm.
|
|
26
|
-
|
|
26
|
+
from chATLAS_Chains.llm.model_selection import (
|
|
27
|
+
GROQ_PRODUCTION_MODELS,
|
|
28
|
+
ChatModelKwargs,
|
|
29
|
+
get_chat_model,
|
|
30
|
+
sanitize_chat_model_kwargs,
|
|
31
|
+
)
|
|
27
32
|
from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
|
|
28
33
|
from chATLAS_Chains.query.query_rewriting import rewrite_query
|
|
29
34
|
from chATLAS_Chains.search.basic import search_runnable
|
|
30
35
|
from chATLAS_Chains.utils.doc_utils import combine_documents
|
|
31
|
-
from chATLAS_Embed.
|
|
36
|
+
from chATLAS_Embed.VectorStores import VectorStore
|
|
37
|
+
|
|
38
|
+
logger = logging.getLogger(__name__)
|
|
32
39
|
|
|
33
40
|
|
|
34
41
|
# Define TypedDict for the simplified state
|
|
@@ -37,26 +44,24 @@ class HybridGraphState(TypedDict, total=False):
|
|
|
37
44
|
search_kwargs: dict
|
|
38
45
|
docs: list[Document]
|
|
39
46
|
answer: str
|
|
47
|
+
chat_history: str # Optional
|
|
40
48
|
|
|
41
49
|
|
|
42
50
|
def advanced_rag(
|
|
43
51
|
vectorstore: VectorStore | list[VectorStore],
|
|
44
52
|
model_name: str,
|
|
45
53
|
prompt: str | None = None,
|
|
46
|
-
|
|
47
|
-
temperature: float = 0.1,
|
|
48
|
-
use_preview_models: bool = False,
|
|
54
|
+
chat_model_kwargs: ChatModelKwargs | None = None,
|
|
49
55
|
enable_query_rewriting: bool = False,
|
|
50
56
|
enable_rrf: bool = False,
|
|
51
57
|
enable_reranking: bool = False,
|
|
52
58
|
# enable_self_evaluation: bool = False,
|
|
53
59
|
query_rewriting_model: str = GROQ_PRODUCTION_MODELS[0],
|
|
54
|
-
|
|
60
|
+
query_rewriting_chat_model_kwargs: ChatModelKwargs | None = None,
|
|
55
61
|
rerank_model: str = "cohere-rerank-3.5",
|
|
56
62
|
pinecone_api_key: str | None = None,
|
|
57
63
|
rrf_constant: float = 60.0,
|
|
58
64
|
rrf_weights: dict[str, float] | None = None,
|
|
59
|
-
retry_config: RetryConfig | None = None,
|
|
60
65
|
) -> CompiledStateGraph:
|
|
61
66
|
"""
|
|
62
67
|
Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
|
|
@@ -64,19 +69,17 @@ def advanced_rag(
|
|
|
64
69
|
:param prompt: The prompt template to use for the language model. If None, uses chATLAS_Chains.prompt.starters.CHAT_PROMPT_TEMPLATE
|
|
65
70
|
:param vectorstore: Single vectorstore instance or list of vectorstore instances to search
|
|
66
71
|
:param model_name: The name of the language model to use for generating responses.
|
|
67
|
-
:param
|
|
68
|
-
:param temperature: Temperature to use for the model.
|
|
69
|
-
:param use_preview_models: If True, allows the use of preview models from Groq.
|
|
72
|
+
:param chat_model_kwargs: Optional kwargs passed through to ``get_chat_model`` for the main model.
|
|
70
73
|
:param enable_query_rewriting: Whether to enable LLM-powered query rewriting.
|
|
71
74
|
:param enable_rrf: Whether to enable RRF (Reciprocal Rank Fusion) for combining results from multiple vectorstores.
|
|
72
75
|
:param enable_reranking: Whether to rerank the retrieved results using the Pinecone API.
|
|
73
76
|
:param query_rewriting_model: Name of the model to use for query rewriting.
|
|
74
|
-
:param
|
|
77
|
+
:param query_rewriting_chat_model_kwargs: Optional kwargs for the query rewriting model.
|
|
78
|
+
Inheritance order is: ``chat_model_kwargs`` -> ``{"temperature": 0.1}`` defaults -> explicit query rewriting kwargs.
|
|
75
79
|
:param rerank_model: Name of the Pinecone reranker model to use.
|
|
76
80
|
:param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
|
|
77
81
|
:param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
|
|
78
82
|
:param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
|
|
79
|
-
:param retry_config: Optional RetryConfig for Groq API calls
|
|
80
83
|
|
|
81
84
|
:return: A compiled LangGraph with the chosen features
|
|
82
85
|
|
|
@@ -90,15 +93,18 @@ def advanced_rag(
|
|
|
90
93
|
# Parallel searcher for vectorstore(s)
|
|
91
94
|
searcher = search_runnable(vectorstore)
|
|
92
95
|
|
|
96
|
+
primary_model_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
|
|
97
|
+
|
|
93
98
|
# Initialise models
|
|
94
|
-
model = get_chat_model(model_name,
|
|
99
|
+
model = get_chat_model(model_name, **primary_model_kwargs)
|
|
95
100
|
|
|
96
101
|
if enable_query_rewriting:
|
|
97
|
-
|
|
98
|
-
|
|
99
|
-
temperature
|
|
100
|
-
|
|
101
|
-
)
|
|
102
|
+
query_model_kwargs: ChatModelKwargs = {
|
|
103
|
+
**primary_model_kwargs,
|
|
104
|
+
"temperature": 0.1,
|
|
105
|
+
}
|
|
106
|
+
query_model_kwargs.update(sanitize_chat_model_kwargs(query_rewriting_chat_model_kwargs))
|
|
107
|
+
query_rewriting_model_instance = get_chat_model(model_name=query_rewriting_model, **query_model_kwargs)
|
|
102
108
|
|
|
103
109
|
# --------------- Define functions for the graph nodes ---------------
|
|
104
110
|
|
|
@@ -107,7 +113,7 @@ def advanced_rag(
|
|
|
107
113
|
question = state["question"]
|
|
108
114
|
|
|
109
115
|
if not enable_query_rewriting:
|
|
110
|
-
|
|
116
|
+
logger.warning("Query rewriting disabled, but query_rewriting node was called.")
|
|
111
117
|
# return original state
|
|
112
118
|
return {**state}
|
|
113
119
|
|
|
@@ -128,7 +134,7 @@ def advanced_rag(
|
|
|
128
134
|
|
|
129
135
|
search_kwargs = state.get("search_kwargs", {})
|
|
130
136
|
if not search_kwargs:
|
|
131
|
-
|
|
137
|
+
logger.warning("No search_kwargs provided, using defaults.")
|
|
132
138
|
results = searcher.invoke(question)
|
|
133
139
|
else:
|
|
134
140
|
results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
|
|
@@ -136,7 +142,7 @@ def advanced_rag(
|
|
|
136
142
|
if "docs" not in results:
|
|
137
143
|
raise Exception("Search results missing 'docs' field")
|
|
138
144
|
|
|
139
|
-
|
|
145
|
+
logger.debug(f"Retrieved {len(results['docs'])} documents")
|
|
140
146
|
|
|
141
147
|
return {
|
|
142
148
|
**state,
|
|
@@ -150,7 +156,7 @@ def advanced_rag(
|
|
|
150
156
|
"""
|
|
151
157
|
|
|
152
158
|
if not enable_reranking:
|
|
153
|
-
|
|
159
|
+
logger.warning("Reranking disabled, but rerank node was called.")
|
|
154
160
|
# return original state
|
|
155
161
|
return {**state}
|
|
156
162
|
|
|
@@ -166,7 +172,7 @@ def advanced_rag(
|
|
|
166
172
|
)
|
|
167
173
|
|
|
168
174
|
except Exception as e:
|
|
169
|
-
|
|
175
|
+
logger.warning(f"Reranking failed with exception {e}. Returning original documents.")
|
|
170
176
|
reranked_docs = docs
|
|
171
177
|
|
|
172
178
|
return {
|
|
@@ -177,13 +183,13 @@ def advanced_rag(
|
|
|
177
183
|
def rrf(state: HybridGraphState) -> HybridGraphState:
|
|
178
184
|
"""Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
|
|
179
185
|
if not enable_rrf:
|
|
180
|
-
|
|
186
|
+
logger.warning("RRF disabled, but rrf node was called.")
|
|
181
187
|
# return original state
|
|
182
188
|
return {**state}
|
|
183
189
|
|
|
184
190
|
docs = state.get("docs", [])
|
|
185
191
|
if not docs:
|
|
186
|
-
|
|
192
|
+
logger.warning("No documents retrieved for RRF.")
|
|
187
193
|
return {**state}
|
|
188
194
|
|
|
189
195
|
try:
|
|
@@ -191,7 +197,7 @@ def advanced_rag(
|
|
|
191
197
|
results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
|
|
192
198
|
)
|
|
193
199
|
except Exception as e:
|
|
194
|
-
|
|
200
|
+
logger.warning(f"RRF failed with exception {e}. Returning original documents.")
|
|
195
201
|
rrf_docs = docs
|
|
196
202
|
|
|
197
203
|
return {
|
|
@@ -201,9 +207,13 @@ def advanced_rag(
|
|
|
201
207
|
|
|
202
208
|
def generate_answer(state: HybridGraphState) -> HybridGraphState:
|
|
203
209
|
"""Generate a response from the LLM"""
|
|
204
|
-
|
|
205
210
|
# Format the prompt using LangChain template
|
|
206
|
-
prompt_input = {
|
|
211
|
+
prompt_input = {
|
|
212
|
+
"context": combine_documents(state["docs"]),
|
|
213
|
+
"question": state["question"],
|
|
214
|
+
"chat_history": state.get("chat_history", ""),
|
|
215
|
+
}
|
|
216
|
+
|
|
207
217
|
final_prompt = prompt_template.format_messages(**prompt_input)
|
|
208
218
|
|
|
209
219
|
response = model.invoke(final_prompt)
|
|
@@ -266,42 +276,28 @@ if __name__ == "__main__":
|
|
|
266
276
|
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
267
277
|
|
|
268
278
|
twiki = get_vectorstore("twiki_prod")
|
|
269
|
-
|
|
270
|
-
retry_config = RetryConfig(
|
|
271
|
-
max_retries=1,
|
|
272
|
-
max_delay=120.0,
|
|
273
|
-
)
|
|
279
|
+
mkdocs = get_vectorstore("mkdocs_prod_v1")
|
|
274
280
|
|
|
275
281
|
# Create the hybrid graph
|
|
276
282
|
graph = advanced_rag(
|
|
277
|
-
vectorstore=[twiki],
|
|
283
|
+
vectorstore=[twiki, mkdocs],
|
|
278
284
|
model_name=GROQ_PRODUCTION_MODELS[0],
|
|
279
285
|
enable_query_rewriting=True,
|
|
280
286
|
enable_rrf=True,
|
|
281
|
-
enable_reranking=
|
|
287
|
+
enable_reranking=False,
|
|
288
|
+
)
|
|
289
|
+
|
|
290
|
+
ans = graph.invoke(
|
|
291
|
+
{
|
|
292
|
+
"question": "What is the crack veto in electron reconstruction?",
|
|
293
|
+
"search_kwargs": {
|
|
294
|
+
"k_text": 3,
|
|
295
|
+
"k": 15,
|
|
296
|
+
"date_filter": "01-01-2010",
|
|
297
|
+
# "type": ["twiki"],
|
|
298
|
+
},
|
|
299
|
+
}
|
|
282
300
|
)
|
|
283
301
|
|
|
284
|
-
|
|
285
|
-
|
|
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)
|
|
302
|
+
print(f"Number of docs is : {len(ans['docs'])}")
|
|
303
|
+
print(f"Answer: {ans['answer']}")
|
|
@@ -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
|
{
|