AutoRAG 0.0.0__py3-none-any.whl
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.
- autorag/__init__.py +82 -0
- autorag/chunker.py +51 -0
- autorag/cli.py +209 -0
- autorag/dashboard.py +199 -0
- autorag/data/__init__.py +109 -0
- autorag/data/chunk/__init__.py +2 -0
- autorag/data/chunk/base.py +128 -0
- autorag/data/chunk/langchain_chunk.py +76 -0
- autorag/data/chunk/llama_index_chunk.py +96 -0
- autorag/data/chunk/run.py +38 -0
- autorag/data/legacy/__init__.py +0 -0
- autorag/data/legacy/corpus/__init__.py +2 -0
- autorag/data/legacy/corpus/langchain.py +47 -0
- autorag/data/legacy/corpus/llama_index.py +93 -0
- autorag/data/legacy/qacreation/__init__.py +6 -0
- autorag/data/legacy/qacreation/base.py +239 -0
- autorag/data/legacy/qacreation/llama_index.py +253 -0
- autorag/data/legacy/qacreation/llama_index_default_prompt.txt +54 -0
- autorag/data/legacy/qacreation/ragas.py +75 -0
- autorag/data/legacy/qacreation/simple.py +99 -0
- autorag/data/parse/__init__.py +1 -0
- autorag/data/parse/base.py +79 -0
- autorag/data/parse/clova.py +194 -0
- autorag/data/parse/langchain_parse.py +87 -0
- autorag/data/parse/llamaparse.py +126 -0
- autorag/data/parse/run.py +141 -0
- autorag/data/parse/table_hybrid_parse.py +134 -0
- autorag/data/qa/__init__.py +3 -0
- autorag/data/qa/evolve/__init__.py +0 -0
- autorag/data/qa/evolve/llama_index_query_evolve.py +64 -0
- autorag/data/qa/evolve/openai_query_evolve.py +81 -0
- autorag/data/qa/evolve/prompt.py +288 -0
- autorag/data/qa/extract_evidence.py +1 -0
- autorag/data/qa/filter/__init__.py +0 -0
- autorag/data/qa/filter/dontknow.py +117 -0
- autorag/data/qa/filter/passage_dependency.py +88 -0
- autorag/data/qa/filter/prompt.py +73 -0
- autorag/data/qa/generation_gt/__init__.py +0 -0
- autorag/data/qa/generation_gt/base.py +16 -0
- autorag/data/qa/generation_gt/llama_index_gen_gt.py +41 -0
- autorag/data/qa/generation_gt/openai_gen_gt.py +84 -0
- autorag/data/qa/generation_gt/prompt.py +27 -0
- autorag/data/qa/query/__init__.py +0 -0
- autorag/data/qa/query/llama_gen_query.py +82 -0
- autorag/data/qa/query/openai_gen_query.py +95 -0
- autorag/data/qa/query/prompt.py +201 -0
- autorag/data/qa/sample.py +26 -0
- autorag/data/qa/schema.py +322 -0
- autorag/data/utils/__init__.py +0 -0
- autorag/data/utils/util.py +103 -0
- autorag/deploy/__init__.py +9 -0
- autorag/deploy/api.py +303 -0
- autorag/deploy/base.py +235 -0
- autorag/deploy/gradio.py +74 -0
- autorag/deploy/swagger.yml +202 -0
- autorag/embedding/__init__.py +0 -0
- autorag/embedding/base.py +144 -0
- autorag/embedding/vllm.py +256 -0
- autorag/evaluation/__init__.py +3 -0
- autorag/evaluation/generation.py +88 -0
- autorag/evaluation/metric/__init__.py +22 -0
- autorag/evaluation/metric/deepeval_prompt.py +322 -0
- autorag/evaluation/metric/g_eval_prompts/coh_detailed.txt +32 -0
- autorag/evaluation/metric/g_eval_prompts/con_detailed.txt +33 -0
- autorag/evaluation/metric/g_eval_prompts/flu_detailed.txt +26 -0
- autorag/evaluation/metric/g_eval_prompts/rel_detailed.txt +33 -0
- autorag/evaluation/metric/generation.py +504 -0
- autorag/evaluation/metric/retrieval.py +115 -0
- autorag/evaluation/metric/retrieval_contents.py +65 -0
- autorag/evaluation/metric/util.py +88 -0
- autorag/evaluation/retrieval.py +83 -0
- autorag/evaluation/retrieval_contents.py +65 -0
- autorag/evaluation/util.py +43 -0
- autorag/evaluator.py +559 -0
- autorag/node_line.py +65 -0
- autorag/nodes/__init__.py +0 -0
- autorag/nodes/generator/__init__.py +4 -0
- autorag/nodes/generator/base.py +103 -0
- autorag/nodes/generator/llama_index_llm.py +169 -0
- autorag/nodes/generator/openai_llm.py +329 -0
- autorag/nodes/generator/run.py +148 -0
- autorag/nodes/generator/vllm.py +147 -0
- autorag/nodes/generator/vllm_api.py +191 -0
- autorag/nodes/hybridretrieval/__init__.py +2 -0
- autorag/nodes/hybridretrieval/base.py +58 -0
- autorag/nodes/hybridretrieval/hybrid_cc.py +227 -0
- autorag/nodes/hybridretrieval/hybrid_rrf.py +149 -0
- autorag/nodes/hybridretrieval/run.py +137 -0
- autorag/nodes/lexicalretrieval/__init__.py +1 -0
- autorag/nodes/lexicalretrieval/bm25.py +381 -0
- autorag/nodes/lexicalretrieval/run.py +148 -0
- autorag/nodes/passageaugmenter/__init__.py +2 -0
- autorag/nodes/passageaugmenter/base.py +76 -0
- autorag/nodes/passageaugmenter/pass_passage_augmenter.py +43 -0
- autorag/nodes/passageaugmenter/prev_next_augmenter.py +155 -0
- autorag/nodes/passageaugmenter/run.py +131 -0
- autorag/nodes/passagecompressor/__init__.py +4 -0
- autorag/nodes/passagecompressor/base.py +78 -0
- autorag/nodes/passagecompressor/longllmlingua.py +115 -0
- autorag/nodes/passagecompressor/pass_compressor.py +16 -0
- autorag/nodes/passagecompressor/refine.py +54 -0
- autorag/nodes/passagecompressor/run.py +186 -0
- autorag/nodes/passagecompressor/tree_summarize.py +56 -0
- autorag/nodes/passagefilter/__init__.py +6 -0
- autorag/nodes/passagefilter/base.py +40 -0
- autorag/nodes/passagefilter/pass_passage_filter.py +14 -0
- autorag/nodes/passagefilter/percentile_cutoff.py +58 -0
- autorag/nodes/passagefilter/recency.py +105 -0
- autorag/nodes/passagefilter/run.py +138 -0
- autorag/nodes/passagefilter/similarity_percentile_cutoff.py +134 -0
- autorag/nodes/passagefilter/similarity_threshold_cutoff.py +112 -0
- autorag/nodes/passagefilter/threshold_cutoff.py +78 -0
- autorag/nodes/passagereranker/__init__.py +16 -0
- autorag/nodes/passagereranker/base.py +44 -0
- autorag/nodes/passagereranker/cohere.py +118 -0
- autorag/nodes/passagereranker/colbert.py +213 -0
- autorag/nodes/passagereranker/flag_embedding.py +112 -0
- autorag/nodes/passagereranker/flag_embedding_llm.py +101 -0
- autorag/nodes/passagereranker/flashrank.py +245 -0
- autorag/nodes/passagereranker/jina.py +115 -0
- autorag/nodes/passagereranker/koreranker.py +136 -0
- autorag/nodes/passagereranker/mixedbreadai.py +126 -0
- autorag/nodes/passagereranker/monot5.py +190 -0
- autorag/nodes/passagereranker/openvino.py +191 -0
- autorag/nodes/passagereranker/pass_reranker.py +31 -0
- autorag/nodes/passagereranker/rankgpt.py +170 -0
- autorag/nodes/passagereranker/run.py +145 -0
- autorag/nodes/passagereranker/sentence_transformer.py +129 -0
- autorag/nodes/passagereranker/tart/__init__.py +1 -0
- autorag/nodes/passagereranker/tart/modeling_enc_t5.py +152 -0
- autorag/nodes/passagereranker/tart/tart.py +139 -0
- autorag/nodes/passagereranker/tart/tokenization_enc_t5.py +112 -0
- autorag/nodes/passagereranker/time_reranker.py +72 -0
- autorag/nodes/passagereranker/upr.py +160 -0
- autorag/nodes/passagereranker/voyageai.py +109 -0
- autorag/nodes/promptmaker/__init__.py +12 -0
- autorag/nodes/promptmaker/base.py +32 -0
- autorag/nodes/promptmaker/chat_fstring.py +73 -0
- autorag/nodes/promptmaker/fstring.py +49 -0
- autorag/nodes/promptmaker/long_context_reorder.py +83 -0
- autorag/nodes/promptmaker/run.py +283 -0
- autorag/nodes/promptmaker/window_replacement.py +85 -0
- autorag/nodes/queryexpansion/__init__.py +4 -0
- autorag/nodes/queryexpansion/base.py +62 -0
- autorag/nodes/queryexpansion/hyde.py +43 -0
- autorag/nodes/queryexpansion/multi_query_expansion.py +57 -0
- autorag/nodes/queryexpansion/pass_query_expansion.py +22 -0
- autorag/nodes/queryexpansion/query_decompose.py +111 -0
- autorag/nodes/queryexpansion/run.py +308 -0
- autorag/nodes/retrieval/__init__.py +0 -0
- autorag/nodes/retrieval/base.py +127 -0
- autorag/nodes/retrieval/run_util.py +152 -0
- autorag/nodes/semanticretrieval/__init__.py +1 -0
- autorag/nodes/semanticretrieval/run.py +148 -0
- autorag/nodes/semanticretrieval/vectordb.py +339 -0
- autorag/nodes/util.py +16 -0
- autorag/parser.py +37 -0
- autorag/schema/__init__.py +3 -0
- autorag/schema/base.py +35 -0
- autorag/schema/metricinput.py +99 -0
- autorag/schema/module.py +24 -0
- autorag/schema/node.py +144 -0
- autorag/strategy.py +165 -0
- autorag/support.py +235 -0
- autorag/utils/__init__.py +8 -0
- autorag/utils/cast.py +45 -0
- autorag/utils/preprocess.py +149 -0
- autorag/utils/util.py +759 -0
- autorag/validator.py +98 -0
- autorag/vectordb/__init__.py +75 -0
- autorag/vectordb/base.py +73 -0
- autorag/vectordb/chroma.py +118 -0
- autorag/vectordb/couchbase.py +239 -0
- autorag/vectordb/milvus.py +169 -0
- autorag/vectordb/pinecone.py +121 -0
- autorag/vectordb/qdrant.py +155 -0
- autorag/vectordb/weaviate.py +184 -0
- autorag/web.py +81 -0
- autorag-0.0.0.dist-info/METADATA +780 -0
- autorag-0.0.0.dist-info/RECORD +184 -0
- autorag-0.0.0.dist-info/WHEEL +5 -0
- autorag-0.0.0.dist-info/entry_points.txt +2 -0
- autorag-0.0.0.dist-info/licenses/LICENSE +201 -0
- autorag-0.0.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,169 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
from typing import Any, Dict, List, Tuple, Optional, Union
|
|
3
|
+
|
|
4
|
+
from pymilvus import (
|
|
5
|
+
DataType,
|
|
6
|
+
FieldSchema,
|
|
7
|
+
CollectionSchema,
|
|
8
|
+
connections,
|
|
9
|
+
Collection,
|
|
10
|
+
MilvusException,
|
|
11
|
+
)
|
|
12
|
+
from pymilvus.orm import utility
|
|
13
|
+
|
|
14
|
+
from autorag.utils.util import apply_recursive
|
|
15
|
+
from autorag.vectordb import BaseVectorStore
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
logger = logging.getLogger("AutoRAG")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class Milvus(BaseVectorStore):
|
|
22
|
+
def __init__(
|
|
23
|
+
self,
|
|
24
|
+
embedding_model: Union[str, List[dict]],
|
|
25
|
+
collection_name: str,
|
|
26
|
+
embedding_batch: int = 100,
|
|
27
|
+
similarity_metric: str = "cosine",
|
|
28
|
+
index_type: str = "IVF_FLAT",
|
|
29
|
+
uri: str = "http://localhost:19530",
|
|
30
|
+
db_name: str = "",
|
|
31
|
+
token: str = "",
|
|
32
|
+
user: str = "",
|
|
33
|
+
password: str = "",
|
|
34
|
+
timeout: Optional[float] = None,
|
|
35
|
+
params: Dict[str, Any] = {},
|
|
36
|
+
):
|
|
37
|
+
super().__init__(embedding_model, similarity_metric, embedding_batch)
|
|
38
|
+
|
|
39
|
+
# Connect to Milvus server
|
|
40
|
+
connections.connect(
|
|
41
|
+
"default",
|
|
42
|
+
uri=uri,
|
|
43
|
+
token=token,
|
|
44
|
+
db_name=db_name,
|
|
45
|
+
user=user,
|
|
46
|
+
password=password,
|
|
47
|
+
)
|
|
48
|
+
self.collection_name = collection_name
|
|
49
|
+
self.timeout = timeout
|
|
50
|
+
self.params = params
|
|
51
|
+
self.index_type = index_type
|
|
52
|
+
|
|
53
|
+
# Set Collection
|
|
54
|
+
if not utility.has_collection(collection_name, timeout=timeout):
|
|
55
|
+
# Get the dimension of the embeddings
|
|
56
|
+
test_embedding_result: List[float] = self.embedding.get_query_embedding(
|
|
57
|
+
"test"
|
|
58
|
+
)
|
|
59
|
+
dimension = len(test_embedding_result)
|
|
60
|
+
|
|
61
|
+
pk = FieldSchema(
|
|
62
|
+
name="id",
|
|
63
|
+
dtype=DataType.VARCHAR,
|
|
64
|
+
max_length=128,
|
|
65
|
+
is_primary=True,
|
|
66
|
+
auto_id=False,
|
|
67
|
+
)
|
|
68
|
+
field = FieldSchema(
|
|
69
|
+
name="vector", dtype=DataType.FLOAT_VECTOR, dim=dimension
|
|
70
|
+
)
|
|
71
|
+
schema = CollectionSchema(fields=[pk, field])
|
|
72
|
+
|
|
73
|
+
self.collection = Collection(name=self.collection_name, schema=schema)
|
|
74
|
+
index_params = {
|
|
75
|
+
"metric_type": self.similarity_metric.upper(),
|
|
76
|
+
"index_type": self.index_type.upper(),
|
|
77
|
+
"params": self.params,
|
|
78
|
+
}
|
|
79
|
+
self.collection.create_index(
|
|
80
|
+
field_name="vector", index_params=index_params, timeout=self.timeout
|
|
81
|
+
)
|
|
82
|
+
else:
|
|
83
|
+
self.collection = Collection(name=self.collection_name)
|
|
84
|
+
|
|
85
|
+
async def add(self, ids: List[str], texts: List[str]):
|
|
86
|
+
texts = self.truncated_inputs(texts)
|
|
87
|
+
text_embeddings: List[
|
|
88
|
+
List[float]
|
|
89
|
+
] = await self.embedding.aget_text_embedding_batch(texts)
|
|
90
|
+
self.add_embedding(ids, text_embeddings)
|
|
91
|
+
|
|
92
|
+
def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
|
|
93
|
+
data = list(
|
|
94
|
+
map(lambda _id, vector: {"id": _id, "vector": vector}, ids, embeddings)
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
# Insert data into the collection
|
|
98
|
+
res = self.collection.insert(data=data, timeout=self.timeout)
|
|
99
|
+
assert res.insert_count == len(ids), (
|
|
100
|
+
f"Insertion failed. Try to insert {len(ids)} but only {res['insert_count']} inserted."
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
self.collection.flush(timeout=self.timeout)
|
|
104
|
+
|
|
105
|
+
async def query(
|
|
106
|
+
self, queries: List[str], top_k: int, **kwargs
|
|
107
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
108
|
+
queries = self.truncated_inputs(queries)
|
|
109
|
+
query_embeddings: List[
|
|
110
|
+
List[float]
|
|
111
|
+
] = await self.embedding.aget_text_embedding_batch(queries)
|
|
112
|
+
|
|
113
|
+
self.collection.load(timeout=self.timeout)
|
|
114
|
+
|
|
115
|
+
# Perform similarity search
|
|
116
|
+
results = self.collection.search(
|
|
117
|
+
data=query_embeddings,
|
|
118
|
+
limit=top_k,
|
|
119
|
+
anns_field="vector",
|
|
120
|
+
param={"metric_type": self.similarity_metric.upper()},
|
|
121
|
+
timeout=self.timeout,
|
|
122
|
+
**kwargs,
|
|
123
|
+
)
|
|
124
|
+
|
|
125
|
+
# Extract IDs and distances
|
|
126
|
+
ids = [[str(hit.id) for hit in result] for result in results]
|
|
127
|
+
distances = [[hit.distance for hit in result] for result in results]
|
|
128
|
+
|
|
129
|
+
if self.similarity_metric in ["l2"]:
|
|
130
|
+
distances = apply_recursive(lambda x: -x, distances)
|
|
131
|
+
|
|
132
|
+
return ids, distances
|
|
133
|
+
|
|
134
|
+
async def fetch(self, ids: List[str]) -> List[List[float]]:
|
|
135
|
+
try:
|
|
136
|
+
self.collection.load(timeout=self.timeout)
|
|
137
|
+
except MilvusException as e:
|
|
138
|
+
logger.warning(f"Failed to load collection: {e}")
|
|
139
|
+
return [[]] * len(ids)
|
|
140
|
+
# Fetch vectors by IDs
|
|
141
|
+
results = self.collection.query(
|
|
142
|
+
expr=f"id in {ids}", output_fields=["id", "vector"], timeout=self.timeout
|
|
143
|
+
)
|
|
144
|
+
id_vector_dict = {str(result["id"]): result["vector"] for result in results}
|
|
145
|
+
result = [id_vector_dict[_id] for _id in ids]
|
|
146
|
+
return result
|
|
147
|
+
|
|
148
|
+
async def is_exist(self, ids: List[str]) -> List[bool]:
|
|
149
|
+
try:
|
|
150
|
+
self.collection.load(timeout=self.timeout)
|
|
151
|
+
except MilvusException:
|
|
152
|
+
return [False] * len(ids)
|
|
153
|
+
# Check the existence of IDs
|
|
154
|
+
results = self.collection.query(
|
|
155
|
+
expr=f"id in {ids}", output_fields=["id"], timeout=self.timeout
|
|
156
|
+
)
|
|
157
|
+
# Determine existence
|
|
158
|
+
existing_ids = {str(result["id"]) for result in results}
|
|
159
|
+
return [str(_id) in existing_ids for _id in ids]
|
|
160
|
+
|
|
161
|
+
async def delete(self, ids: List[str]):
|
|
162
|
+
# Delete entries by IDs
|
|
163
|
+
self.collection.delete(expr=f"id in {ids}", timeout=self.timeout)
|
|
164
|
+
|
|
165
|
+
def delete_collection(self):
|
|
166
|
+
# Delete the collection
|
|
167
|
+
self.collection.release(timeout=self.timeout)
|
|
168
|
+
self.collection.drop_index(timeout=self.timeout)
|
|
169
|
+
self.collection.drop(timeout=self.timeout)
|
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
from pinecone.grpc import PineconeGRPC as Pinecone_client
|
|
4
|
+
from pinecone import ServerlessSpec
|
|
5
|
+
|
|
6
|
+
from typing import List, Optional, Tuple, Union
|
|
7
|
+
|
|
8
|
+
from autorag.utils.util import make_batch, apply_recursive
|
|
9
|
+
from autorag.vectordb import BaseVectorStore
|
|
10
|
+
|
|
11
|
+
logger = logging.getLogger("AutoRAG")
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class Pinecone(BaseVectorStore):
|
|
15
|
+
def __init__(
|
|
16
|
+
self,
|
|
17
|
+
embedding_model: Union[str, List[dict]],
|
|
18
|
+
index_name: str,
|
|
19
|
+
embedding_batch: int = 100,
|
|
20
|
+
dimension: int = 1536,
|
|
21
|
+
similarity_metric: str = "cosine", # "cosine", "dotproduct", "euclidean"
|
|
22
|
+
cloud: Optional[str] = "aws",
|
|
23
|
+
region: Optional[str] = "us-east-1",
|
|
24
|
+
api_key: Optional[str] = None,
|
|
25
|
+
deletion_protection: Optional[str] = "disabled", # "enabled" or "disabled"
|
|
26
|
+
namespace: Optional[str] = "default",
|
|
27
|
+
ingest_batch: int = 200,
|
|
28
|
+
):
|
|
29
|
+
super().__init__(embedding_model, similarity_metric, embedding_batch)
|
|
30
|
+
|
|
31
|
+
self.index_name = index_name
|
|
32
|
+
self.namespace = namespace
|
|
33
|
+
self.ingest_batch = ingest_batch
|
|
34
|
+
|
|
35
|
+
self.client = Pinecone_client(api_key=api_key)
|
|
36
|
+
|
|
37
|
+
if similarity_metric == "ip":
|
|
38
|
+
similarity_metric = "dotproduct"
|
|
39
|
+
elif similarity_metric == "l2":
|
|
40
|
+
similarity_metric = "euclidean"
|
|
41
|
+
|
|
42
|
+
if not self.client.has_index(index_name):
|
|
43
|
+
self.client.create_index(
|
|
44
|
+
name=index_name,
|
|
45
|
+
dimension=dimension,
|
|
46
|
+
metric=similarity_metric,
|
|
47
|
+
spec=ServerlessSpec(
|
|
48
|
+
cloud=cloud,
|
|
49
|
+
region=region,
|
|
50
|
+
),
|
|
51
|
+
deletion_protection=deletion_protection,
|
|
52
|
+
)
|
|
53
|
+
self.index = self.client.Index(index_name)
|
|
54
|
+
|
|
55
|
+
async def add(self, ids: List[str], texts: List[str]):
|
|
56
|
+
texts = self.truncated_inputs(texts)
|
|
57
|
+
text_embeddings: List[
|
|
58
|
+
List[float]
|
|
59
|
+
] = await self.embedding.aget_text_embedding_batch(texts)
|
|
60
|
+
self.add_embedding(ids, text_embeddings)
|
|
61
|
+
|
|
62
|
+
def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
|
|
63
|
+
vector_tuples = list(zip(ids, embeddings))
|
|
64
|
+
batch_vectors = make_batch(vector_tuples, self.ingest_batch)
|
|
65
|
+
|
|
66
|
+
async_res = [
|
|
67
|
+
self.index.upsert(
|
|
68
|
+
vectors=batch_vector_tuples,
|
|
69
|
+
namespace=self.namespace,
|
|
70
|
+
async_req=True,
|
|
71
|
+
)
|
|
72
|
+
for batch_vector_tuples in batch_vectors
|
|
73
|
+
]
|
|
74
|
+
# Wait for the async requests to finish
|
|
75
|
+
[async_result.result() for async_result in async_res]
|
|
76
|
+
|
|
77
|
+
async def fetch(self, ids: List[str]) -> List[List[float]]:
|
|
78
|
+
results = self.index.fetch(ids=ids, namespace=self.namespace)
|
|
79
|
+
id_vector_dict = {
|
|
80
|
+
str(key): val["values"] for key, val in results["vectors"].items()
|
|
81
|
+
}
|
|
82
|
+
result = [id_vector_dict[_id] for _id in ids]
|
|
83
|
+
return result
|
|
84
|
+
|
|
85
|
+
async def is_exist(self, ids: List[str]) -> List[bool]:
|
|
86
|
+
fetched_result = self.index.fetch(ids=ids, namespace=self.namespace)
|
|
87
|
+
existed_ids = list(map(str, fetched_result.get("vectors", {}).keys()))
|
|
88
|
+
return list(map(lambda x: x in existed_ids, ids))
|
|
89
|
+
|
|
90
|
+
async def query(
|
|
91
|
+
self, queries: List[str], top_k: int, **kwargs
|
|
92
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
93
|
+
queries = self.truncated_inputs(queries)
|
|
94
|
+
query_embeddings: List[
|
|
95
|
+
List[float]
|
|
96
|
+
] = await self.embedding.aget_text_embedding_batch(queries)
|
|
97
|
+
|
|
98
|
+
ids, scores = [], []
|
|
99
|
+
for query_embedding in query_embeddings:
|
|
100
|
+
response = self.index.query(
|
|
101
|
+
vector=query_embedding,
|
|
102
|
+
top_k=top_k,
|
|
103
|
+
include_values=True,
|
|
104
|
+
namespace=self.namespace,
|
|
105
|
+
)
|
|
106
|
+
|
|
107
|
+
ids.append([o.id for o in response.matches])
|
|
108
|
+
scores.append([o.score for o in response.matches])
|
|
109
|
+
|
|
110
|
+
if self.similarity_metric in ["l2"]:
|
|
111
|
+
scores = apply_recursive(lambda x: -x, scores)
|
|
112
|
+
|
|
113
|
+
return ids, scores
|
|
114
|
+
|
|
115
|
+
async def delete(self, ids: List[str]):
|
|
116
|
+
# Delete entries by IDs
|
|
117
|
+
self.index.delete(ids=ids, namespace=self.namespace)
|
|
118
|
+
|
|
119
|
+
def delete_index(self):
|
|
120
|
+
# Delete the index
|
|
121
|
+
self.client.delete_index(self.index_name)
|
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
from qdrant_client import QdrantClient
|
|
4
|
+
from qdrant_client.models import (
|
|
5
|
+
Distance,
|
|
6
|
+
VectorParams,
|
|
7
|
+
PointStruct,
|
|
8
|
+
PointIdsList,
|
|
9
|
+
HasIdCondition,
|
|
10
|
+
Filter,
|
|
11
|
+
SearchRequest,
|
|
12
|
+
)
|
|
13
|
+
|
|
14
|
+
from typing import List, Tuple, Union
|
|
15
|
+
|
|
16
|
+
from autorag.vectordb import BaseVectorStore
|
|
17
|
+
|
|
18
|
+
logger = logging.getLogger("AutoRAG")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class Qdrant(BaseVectorStore):
|
|
22
|
+
def __init__(
|
|
23
|
+
self,
|
|
24
|
+
embedding_model: Union[str, List[dict]],
|
|
25
|
+
collection_name: str,
|
|
26
|
+
embedding_batch: int = 100,
|
|
27
|
+
similarity_metric: str = "cosine",
|
|
28
|
+
client_type: str = "docker",
|
|
29
|
+
url: str = "http://localhost:6333",
|
|
30
|
+
host: str = "",
|
|
31
|
+
api_key: str = "",
|
|
32
|
+
dimension: int = 1536,
|
|
33
|
+
ingest_batch: int = 64,
|
|
34
|
+
parallel: int = 1,
|
|
35
|
+
max_retries: int = 3,
|
|
36
|
+
):
|
|
37
|
+
super().__init__(embedding_model, similarity_metric, embedding_batch)
|
|
38
|
+
|
|
39
|
+
self.collection_name = collection_name
|
|
40
|
+
self.ingest_batch = ingest_batch
|
|
41
|
+
self.parallel = parallel
|
|
42
|
+
self.max_retries = max_retries
|
|
43
|
+
|
|
44
|
+
if similarity_metric == "cosine":
|
|
45
|
+
distance = Distance.COSINE
|
|
46
|
+
elif similarity_metric == "ip":
|
|
47
|
+
distance = Distance.DOT
|
|
48
|
+
elif similarity_metric == "l2":
|
|
49
|
+
distance = Distance.EUCLID
|
|
50
|
+
else:
|
|
51
|
+
raise ValueError(
|
|
52
|
+
f"similarity_metric {similarity_metric} is not supported\n"
|
|
53
|
+
"supported similarity metrics are: cosine, ip, l2"
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
if client_type == "docker":
|
|
57
|
+
self.client = QdrantClient(
|
|
58
|
+
url=url,
|
|
59
|
+
)
|
|
60
|
+
elif client_type == "cloud":
|
|
61
|
+
self.client = QdrantClient(
|
|
62
|
+
host=host,
|
|
63
|
+
api_key=api_key,
|
|
64
|
+
)
|
|
65
|
+
else:
|
|
66
|
+
raise ValueError(
|
|
67
|
+
f"client_type {client_type} is not supported\n"
|
|
68
|
+
"supported client types are: docker, cloud"
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
if not self.client.collection_exists(collection_name):
|
|
72
|
+
self.client.create_collection(
|
|
73
|
+
collection_name,
|
|
74
|
+
vectors_config=VectorParams(
|
|
75
|
+
size=dimension,
|
|
76
|
+
distance=distance,
|
|
77
|
+
),
|
|
78
|
+
)
|
|
79
|
+
self.collection = self.client.get_collection(collection_name)
|
|
80
|
+
|
|
81
|
+
async def add(self, ids: List[str], texts: List[str]):
|
|
82
|
+
texts = self.truncated_inputs(texts)
|
|
83
|
+
text_embeddings = await self.embedding.aget_text_embedding_batch(texts)
|
|
84
|
+
self.add_embedding(ids, text_embeddings)
|
|
85
|
+
|
|
86
|
+
def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
|
|
87
|
+
points = list(
|
|
88
|
+
map(lambda x: PointStruct(id=x[0], vector=x[1]), zip(ids, embeddings))
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
self.client.upload_points(
|
|
92
|
+
collection_name=self.collection_name,
|
|
93
|
+
points=points,
|
|
94
|
+
batch_size=self.ingest_batch,
|
|
95
|
+
parallel=self.parallel,
|
|
96
|
+
max_retries=self.max_retries,
|
|
97
|
+
wait=True,
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
async def fetch(self, ids: List[str]) -> List[List[float]]:
|
|
101
|
+
# Fetch vectors by IDs
|
|
102
|
+
fetched_results = self.client.retrieve(
|
|
103
|
+
collection_name=self.collection_name,
|
|
104
|
+
ids=ids,
|
|
105
|
+
with_vectors=True,
|
|
106
|
+
)
|
|
107
|
+
return list(map(lambda x: x.vector, fetched_results))
|
|
108
|
+
|
|
109
|
+
async def is_exist(self, ids: List[str]) -> List[bool]:
|
|
110
|
+
existed_result = self.client.scroll(
|
|
111
|
+
collection_name=self.collection_name,
|
|
112
|
+
scroll_filter=Filter(
|
|
113
|
+
must=[
|
|
114
|
+
HasIdCondition(has_id=ids),
|
|
115
|
+
],
|
|
116
|
+
),
|
|
117
|
+
)
|
|
118
|
+
# existed_result is tuple. So we use existed_result[0] to get list of Record
|
|
119
|
+
existed_ids = list(map(lambda x: x.id, existed_result[0]))
|
|
120
|
+
return list(map(lambda x: x in existed_ids, ids))
|
|
121
|
+
|
|
122
|
+
async def query(
|
|
123
|
+
self, queries: List[str], top_k: int, **kwargs
|
|
124
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
125
|
+
queries = self.truncated_inputs(queries)
|
|
126
|
+
query_embeddings: List[
|
|
127
|
+
List[float]
|
|
128
|
+
] = await self.embedding.aget_text_embedding_batch(queries)
|
|
129
|
+
|
|
130
|
+
search_queries = list(
|
|
131
|
+
map(
|
|
132
|
+
lambda x: SearchRequest(vector=x, limit=top_k, with_vector=True),
|
|
133
|
+
query_embeddings,
|
|
134
|
+
)
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
search_result = self.client.search_batch(
|
|
138
|
+
collection_name=self.collection_name, requests=search_queries
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
# Extract IDs and distances
|
|
142
|
+
ids = [[str(hit.id) for hit in result] for result in search_result]
|
|
143
|
+
scores = [[hit.score for hit in result] for result in search_result]
|
|
144
|
+
|
|
145
|
+
return ids, scores
|
|
146
|
+
|
|
147
|
+
async def delete(self, ids: List[str]):
|
|
148
|
+
self.client.delete(
|
|
149
|
+
collection_name=self.collection_name,
|
|
150
|
+
points_selector=PointIdsList(points=ids),
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
def delete_collection(self):
|
|
154
|
+
# Delete the collection
|
|
155
|
+
self.client.delete_collection(self.collection_name)
|
|
@@ -0,0 +1,184 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
import weaviate
|
|
4
|
+
from weaviate.classes.init import Auth
|
|
5
|
+
from weaviate.classes.config import Property, DataType
|
|
6
|
+
import weaviate.classes as wvc
|
|
7
|
+
from weaviate.classes.query import MetadataQuery
|
|
8
|
+
|
|
9
|
+
from typing import List, Optional, Tuple, Union
|
|
10
|
+
|
|
11
|
+
from autorag.vectordb import BaseVectorStore
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger("AutoRAG")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class Weaviate(BaseVectorStore):
|
|
17
|
+
def __init__(
|
|
18
|
+
self,
|
|
19
|
+
embedding_model: Union[str, List[dict]],
|
|
20
|
+
collection_name: str,
|
|
21
|
+
embedding_batch: int = 100,
|
|
22
|
+
similarity_metric: str = "cosine",
|
|
23
|
+
client_type: str = "docker",
|
|
24
|
+
host: str = "localhost",
|
|
25
|
+
port: int = 8080,
|
|
26
|
+
grpc_port: int = 50051,
|
|
27
|
+
url: Optional[str] = None,
|
|
28
|
+
api_key: Optional[str] = None,
|
|
29
|
+
text_key: str = "content",
|
|
30
|
+
):
|
|
31
|
+
super().__init__(embedding_model, similarity_metric, embedding_batch)
|
|
32
|
+
|
|
33
|
+
self.text_key = text_key
|
|
34
|
+
|
|
35
|
+
if client_type == "docker":
|
|
36
|
+
self.client = weaviate.connect_to_local(
|
|
37
|
+
host=host,
|
|
38
|
+
port=port,
|
|
39
|
+
grpc_port=grpc_port,
|
|
40
|
+
)
|
|
41
|
+
elif client_type == "cloud":
|
|
42
|
+
self.client = weaviate.connect_to_weaviate_cloud(
|
|
43
|
+
cluster_url=url,
|
|
44
|
+
auth_credentials=Auth.api_key(api_key),
|
|
45
|
+
)
|
|
46
|
+
else:
|
|
47
|
+
raise ValueError(
|
|
48
|
+
f"client_type {client_type} is not supported\n"
|
|
49
|
+
"supported client types are: docker, cloud"
|
|
50
|
+
)
|
|
51
|
+
if similarity_metric == "cosine":
|
|
52
|
+
distance_metric = wvc.config.VectorDistances.COSINE
|
|
53
|
+
elif similarity_metric == "ip":
|
|
54
|
+
distance_metric = wvc.config.VectorDistances.DOT
|
|
55
|
+
elif similarity_metric == "l2":
|
|
56
|
+
distance_metric = wvc.config.VectorDistances.L2_SQUARED
|
|
57
|
+
else:
|
|
58
|
+
raise ValueError(
|
|
59
|
+
f"similarity_metric {similarity_metric} is not supported\n"
|
|
60
|
+
"supported similarity metrics are: cosine, ip, l2"
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
if not self.client.collections.exists(collection_name):
|
|
64
|
+
self.client.collections.create(
|
|
65
|
+
collection_name,
|
|
66
|
+
properties=[
|
|
67
|
+
Property(
|
|
68
|
+
name="content", data_type=DataType.TEXT, skip_vectorization=True
|
|
69
|
+
),
|
|
70
|
+
],
|
|
71
|
+
vectorizer_config=wvc.config.Configure.Vectorizer.none(),
|
|
72
|
+
vector_index_config=wvc.config.Configure.VectorIndex.hnsw( # hnsw, flat, dynamic,
|
|
73
|
+
distance_metric=distance_metric
|
|
74
|
+
),
|
|
75
|
+
)
|
|
76
|
+
self.collection = self.client.collections.get(collection_name)
|
|
77
|
+
self.collection_name = collection_name
|
|
78
|
+
|
|
79
|
+
async def add(self, ids: List[str], texts: List[str]):
|
|
80
|
+
texts = self.truncated_inputs(texts)
|
|
81
|
+
text_embeddings = await self.embedding.aget_text_embedding_batch(texts)
|
|
82
|
+
|
|
83
|
+
with self.client.batch.dynamic() as batch:
|
|
84
|
+
for i, text in enumerate(texts):
|
|
85
|
+
data_properties = {self.text_key: text}
|
|
86
|
+
|
|
87
|
+
batch.add_object(
|
|
88
|
+
collection=self.collection_name,
|
|
89
|
+
properties=data_properties,
|
|
90
|
+
uuid=ids[i],
|
|
91
|
+
vector=text_embeddings[i],
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
failed_objs = self.client.batch.failed_objects
|
|
95
|
+
for obj in failed_objs:
|
|
96
|
+
err_message = (
|
|
97
|
+
f"Failed to add object: {obj.original_uuid}\nReason: {obj.message}"
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
logger.error(err_message)
|
|
101
|
+
|
|
102
|
+
def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
|
|
103
|
+
with self.client.batch.dynamic() as batch:
|
|
104
|
+
for i in range(len(ids)):
|
|
105
|
+
batch.add_object(
|
|
106
|
+
collection=self.collection_name,
|
|
107
|
+
uuid=ids[i],
|
|
108
|
+
vector=embeddings[i],
|
|
109
|
+
)
|
|
110
|
+
|
|
111
|
+
failed_objs = self.client.batch.failed_objects
|
|
112
|
+
for obj in failed_objs:
|
|
113
|
+
err_message = (
|
|
114
|
+
f"Failed to add object: {obj.original_uuid}\nReason: {obj.message}"
|
|
115
|
+
)
|
|
116
|
+
|
|
117
|
+
logger.error(err_message)
|
|
118
|
+
|
|
119
|
+
async def fetch(self, ids: List[str]) -> List[List[float]]:
|
|
120
|
+
# Fetch vectors by IDs
|
|
121
|
+
results = self.collection.query.fetch_objects(
|
|
122
|
+
filters=wvc.query.Filter.by_property("_id").contains_any(ids),
|
|
123
|
+
include_vector=True,
|
|
124
|
+
)
|
|
125
|
+
id_vector_dict = {
|
|
126
|
+
str(object.uuid): object.vector["default"] for object in results.objects
|
|
127
|
+
}
|
|
128
|
+
result = [id_vector_dict[_id] for _id in ids]
|
|
129
|
+
return result
|
|
130
|
+
|
|
131
|
+
async def is_exist(self, ids: List[str]) -> List[bool]:
|
|
132
|
+
fetched_result = self.collection.query.fetch_objects(
|
|
133
|
+
filters=wvc.query.Filter.by_property("_id").contains_any(ids),
|
|
134
|
+
)
|
|
135
|
+
existed_ids = [str(result.uuid) for result in fetched_result.objects]
|
|
136
|
+
return list(map(lambda x: x in existed_ids, ids))
|
|
137
|
+
|
|
138
|
+
async def query(
|
|
139
|
+
self, queries: List[str], top_k: int, **kwargs
|
|
140
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
141
|
+
queries = self.truncated_inputs(queries)
|
|
142
|
+
query_embeddings: List[
|
|
143
|
+
List[float]
|
|
144
|
+
] = await self.embedding.aget_text_embedding_batch(queries)
|
|
145
|
+
|
|
146
|
+
ids, scores = [], []
|
|
147
|
+
for query_embedding in query_embeddings:
|
|
148
|
+
response = self.collection.query.near_vector(
|
|
149
|
+
near_vector=query_embedding,
|
|
150
|
+
limit=top_k,
|
|
151
|
+
return_metadata=MetadataQuery(distance=True),
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
ids.append([o.uuid for o in response.objects])
|
|
155
|
+
scores.append(
|
|
156
|
+
[
|
|
157
|
+
distance_to_score(o.metadata.distance, self.similarity_metric)
|
|
158
|
+
for o in response.objects
|
|
159
|
+
]
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
return ids, scores
|
|
163
|
+
|
|
164
|
+
async def delete(self, ids: List[str]):
|
|
165
|
+
filter = wvc.query.Filter.by_id().contains_any(ids)
|
|
166
|
+
self.collection.data.delete_many(where=filter)
|
|
167
|
+
|
|
168
|
+
def delete_collection(self):
|
|
169
|
+
# Delete the collection
|
|
170
|
+
self.client.collections.delete(self.collection_name)
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def distance_to_score(distance: float, similarity_metric) -> float:
|
|
174
|
+
if similarity_metric == "cosine":
|
|
175
|
+
return 1 - distance
|
|
176
|
+
elif similarity_metric == "ip":
|
|
177
|
+
return -distance
|
|
178
|
+
elif similarity_metric == "l2":
|
|
179
|
+
return -distance
|
|
180
|
+
else:
|
|
181
|
+
raise ValueError(
|
|
182
|
+
f"similarity_metric {similarity_metric} is not supported\n"
|
|
183
|
+
"supported similarity metrics are: cosine, ip, l2"
|
|
184
|
+
)
|
autorag/web.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import click
|
|
4
|
+
import streamlit as st
|
|
5
|
+
|
|
6
|
+
from autorag.deploy import Runner
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def get_runner(
|
|
10
|
+
yaml_path: Optional[str], project_dir: Optional[str], trial_path: Optional[str]
|
|
11
|
+
):
|
|
12
|
+
if not yaml_path and not trial_path:
|
|
13
|
+
raise ValueError("yaml_path or trial_path must be given.")
|
|
14
|
+
elif yaml_path and trial_path:
|
|
15
|
+
raise ValueError("yaml_path and trial_path cannot be given at the same time.")
|
|
16
|
+
elif yaml_path:
|
|
17
|
+
return Runner.from_yaml(yaml_path, project_dir=project_dir)
|
|
18
|
+
elif trial_path:
|
|
19
|
+
return Runner.from_trial_folder(trial_path)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def set_initial_state():
|
|
23
|
+
if "messages" not in st.session_state:
|
|
24
|
+
st.session_state["messages"] = [
|
|
25
|
+
{
|
|
26
|
+
"role": "assistant",
|
|
27
|
+
"content": "Welcome !",
|
|
28
|
+
}
|
|
29
|
+
]
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def set_page_config():
|
|
33
|
+
st.set_page_config(
|
|
34
|
+
page_title="AutoRAG",
|
|
35
|
+
page_icon="🤖",
|
|
36
|
+
layout="wide",
|
|
37
|
+
initial_sidebar_state="expanded",
|
|
38
|
+
menu_items={
|
|
39
|
+
"Get help": "https://github.com/Marker-Inc-Korea/AutoRAG/discussions",
|
|
40
|
+
"Report a bug": "https://github.com/Marker-Inc-Korea/AutoRAG/issues",
|
|
41
|
+
},
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def set_page_header():
|
|
46
|
+
st.header("📚 AutoRAG", anchor=False)
|
|
47
|
+
st.caption("Input a question and get an answer from the given documents. ")
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def chat_box(runner: Runner):
|
|
51
|
+
if query := st.chat_input("How can I help?"):
|
|
52
|
+
# Add the user input to messages state
|
|
53
|
+
st.session_state["messages"].append({"role": "user", "content": query})
|
|
54
|
+
with st.chat_message("user"):
|
|
55
|
+
st.markdown(query)
|
|
56
|
+
|
|
57
|
+
# Generate llama-index stream with user input
|
|
58
|
+
with st.chat_message("assistant"):
|
|
59
|
+
with st.spinner("Processing..."):
|
|
60
|
+
response = st.write(runner.run(query))
|
|
61
|
+
|
|
62
|
+
# Add the final response to messages state
|
|
63
|
+
st.session_state["messages"].append({"role": "assistant", "content": response})
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
@click.command()
|
|
67
|
+
@click.option("--yaml_path", type=str, help="Path to the YAML file.")
|
|
68
|
+
@click.option("--project_dir", type=str, help="Path to the project directory.")
|
|
69
|
+
@click.option("--trial_path", type=str, help="Path to the trial directory.")
|
|
70
|
+
def run_web_server(
|
|
71
|
+
yaml_path: Optional[str], project_dir: Optional[str], trial_path: Optional[str]
|
|
72
|
+
):
|
|
73
|
+
runner = get_runner(yaml_path, project_dir, trial_path)
|
|
74
|
+
set_initial_state()
|
|
75
|
+
set_page_config()
|
|
76
|
+
set_page_header()
|
|
77
|
+
chat_box(runner)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
if __name__ == "__main__":
|
|
81
|
+
run_web_server()
|