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.
- {chatlas_chains-0.1.3/chATLAS_Chains.egg-info → chatlas_chains-0.1.4}/PKG-INFO +47 -22
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/README.md +37 -17
- chatlas_chains-0.1.4/chATLAS_Chains/chains/advanced.py +299 -0
- chatlas_chains-0.1.4/chATLAS_Chains/chains/basic.py +82 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/chains/basic_graph.py +40 -19
- chatlas_chains-0.1.4/chATLAS_Chains/chains/enhanced_agentic_graph.py +873 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/chains/websearch_retrieval_chain.py +17 -21
- chatlas_chains-0.1.4/chATLAS_Chains/documents/rerank.py +96 -0
- chatlas_chains-0.1.4/chATLAS_Chains/documents/rrf.py +147 -0
- chatlas_chains-0.1.4/chATLAS_Chains/llm/groq.py +598 -0
- chatlas_chains-0.1.4/chATLAS_Chains/llm/model_selection.py +172 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/prompt/starters.py +2 -4
- chatlas_chains-0.1.4/chATLAS_Chains/query/query_rewriting.py +56 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/search/basic.py +22 -10
- chatlas_chains-0.1.4/chATLAS_Chains/utils/doc_utils.py +67 -0
- chatlas_chains-0.1.4/chATLAS_Chains/vectorstore.py +190 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4/chATLAS_Chains.egg-info}/PKG-INFO +47 -22
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains.egg-info/SOURCES.txt +7 -3
- chatlas_chains-0.1.4/chATLAS_Chains.egg-info/requires.txt +11 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains.egg-info/top_level.txt +0 -1
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/pyproject.toml +21 -5
- chatlas_chains-0.1.4/tests/conftest.py +68 -0
- chatlas_chains-0.1.4/tests/test_chains.py +140 -0
- chatlas_chains-0.1.4/tests/test_groq.py +786 -0
- chatlas_chains-0.1.4/tests/test_llm.py +205 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/tests/test_search.py +16 -13
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/tests/test_utils.py +2 -2
- chatlas_chains-0.1.3/benchmark/basic.py +0 -34
- chatlas_chains-0.1.3/benchmark/conversational.py +0 -0
- chatlas_chains-0.1.3/chATLAS_Chains/chains/basic.py +0 -42
- chatlas_chains-0.1.3/chATLAS_Chains/chains/conversational.py +0 -0
- chatlas_chains-0.1.3/chATLAS_Chains/llm/model_selection.py +0 -47
- chatlas_chains-0.1.3/chATLAS_Chains/utils/doc_utils.py +0 -28
- chatlas_chains-0.1.3/chATLAS_Chains/vectorstore.py +0 -94
- chatlas_chains-0.1.3/chATLAS_Chains.egg-info/requires.txt +0 -5
- chatlas_chains-0.1.3/tests/conftest.py +0 -274
- chatlas_chains-0.1.3/tests/test_chains.py +0 -54
- chatlas_chains-0.1.3/tests/test_llm.py +0 -81
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/LICENSE +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/__init__.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/chains/__init__.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/llm/__init__.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/log.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/prompt/__init__.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/search/__init__.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains/utils/__init__.py +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
- {chatlas_chains-0.1.3 → chatlas_chains-0.1.4}/setup.cfg +0 -0
- {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
|
+
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/
|
|
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:
|
|
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
|
-
|
|
31
|
-
|
|
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=
|
|
54
|
+
export CHATLAS_OPENAI_KEY="your api key"
|
|
42
55
|
```
|
|
43
56
|
|
|
44
|
-
2.
|
|
57
|
+
2. Using LLMs via the Groq API
|
|
45
58
|
```bash
|
|
46
|
-
export
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
##
|
|
69
|
+
## Available Chains
|
|
70
|
+
- chains.basic.basic_retrieval_chain
|
|
71
|
+
- chains.basic_graph.basic_retrieval_graph
|
|
72
|
+
- chains.advanced.advanced_rag
|
|
61
73
|
|
|
62
|
-
|
|
74
|
+
## Postgres
|
|
63
75
|
|
|
64
|
-
If you want to create a local
|
|
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
|
-
|
|
10
|
-
|
|
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=
|
|
28
|
+
export CHATLAS_OPENAI_KEY="your api key"
|
|
21
29
|
```
|
|
22
30
|
|
|
23
|
-
2.
|
|
31
|
+
2. Using LLMs via the Groq API
|
|
24
32
|
```bash
|
|
25
|
-
export
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
##
|
|
43
|
+
## Available Chains
|
|
44
|
+
- chains.basic.basic_retrieval_chain
|
|
45
|
+
- chains.basic_graph.basic_retrieval_graph
|
|
46
|
+
- chains.advanced.advanced_rag
|
|
40
47
|
|
|
41
|
-
|
|
48
|
+
## Postgres
|
|
42
49
|
|
|
43
|
-
If you want to create a local
|
|
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(
|
|
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
|
-
|
|
88
|
-
|
|
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
|
|
140
|
-
import
|
|
156
|
+
# Example usage
|
|
157
|
+
from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
|
|
158
|
+
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
141
159
|
|
|
142
|
-
|
|
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
|
-
|
|
147
|
-
|
|
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(
|
|
167
|
+
graph = basic_retrieval_graph(
|
|
168
|
+
prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vs, model_name=model_name, use_preview_models=True
|
|
169
|
+
)
|
|
150
170
|
|
|
151
|
-
|
|
171
|
+
result = graph.invoke(
|
|
152
172
|
{
|
|
153
|
-
"question": "
|
|
173
|
+
"question": "What is the crack veto for electron reconstruction",
|
|
154
174
|
"search_kwargs": {
|
|
155
|
-
"k_text":
|
|
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(
|
|
183
|
+
print("Num retrieved documents: ", len(result["merged_docs"]))
|
|
184
|
+
print("Answer: ", result["answer"])
|