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
autorag/validator.py
ADDED
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
import itertools
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
import tempfile
|
|
5
|
+
|
|
6
|
+
import pandas as pd
|
|
7
|
+
|
|
8
|
+
from autorag.evaluator import Evaluator
|
|
9
|
+
from autorag.utils import (
|
|
10
|
+
cast_qa_dataset,
|
|
11
|
+
cast_corpus_dataset,
|
|
12
|
+
validate_qa_from_corpus_dataset,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
logger = logging.getLogger("AutoRAG")
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class Validator:
|
|
19
|
+
def __init__(self, qa_data_path: str, corpus_data_path: str):
|
|
20
|
+
"""
|
|
21
|
+
Initialize a Validator object.
|
|
22
|
+
|
|
23
|
+
:param qa_data_path: The path to the QA dataset.
|
|
24
|
+
Must be parquet file.
|
|
25
|
+
:param corpus_data_path: The path to the corpus dataset.
|
|
26
|
+
Must be parquet file.
|
|
27
|
+
"""
|
|
28
|
+
# validate data paths
|
|
29
|
+
if not os.path.exists(qa_data_path):
|
|
30
|
+
raise ValueError(f"QA data path {qa_data_path} does not exist.")
|
|
31
|
+
if not os.path.exists(corpus_data_path):
|
|
32
|
+
raise ValueError(f"Corpus data path {corpus_data_path} does not exist.")
|
|
33
|
+
if not qa_data_path.endswith(".parquet"):
|
|
34
|
+
raise ValueError(f"QA data path {qa_data_path} is not a parquet file.")
|
|
35
|
+
if not corpus_data_path.endswith(".parquet"):
|
|
36
|
+
raise ValueError(
|
|
37
|
+
f"Corpus data path {corpus_data_path} is not a parquet file."
|
|
38
|
+
)
|
|
39
|
+
self.qa_data = pd.read_parquet(qa_data_path, engine="pyarrow")
|
|
40
|
+
self.corpus_data = pd.read_parquet(corpus_data_path, engine="pyarrow")
|
|
41
|
+
self.qa_data = cast_qa_dataset(self.qa_data)
|
|
42
|
+
self.corpus_data = cast_corpus_dataset(self.corpus_data)
|
|
43
|
+
|
|
44
|
+
def validate(self, yaml_path: str, qa_cnt: int = 5, random_state: int = 42):
|
|
45
|
+
# Determine the sample size and log a warning if qa_cnt is larger than available records
|
|
46
|
+
available_records = len(self.qa_data)
|
|
47
|
+
safe_sample_size = min(qa_cnt, available_records) # 먼저 safe_sample_size 계산
|
|
48
|
+
|
|
49
|
+
if safe_sample_size < qa_cnt:
|
|
50
|
+
logger.warning(
|
|
51
|
+
f"Minimal Requested sample size ({qa_cnt}) is larger than available records ({available_records}). "
|
|
52
|
+
f"Sampling will be limited to {safe_sample_size} records. "
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
# safe sample QA data
|
|
56
|
+
sample_qa_df = self.qa_data.sample(
|
|
57
|
+
n=safe_sample_size, random_state=random_state
|
|
58
|
+
)
|
|
59
|
+
sample_qa_df.reset_index(drop=True, inplace=True)
|
|
60
|
+
|
|
61
|
+
# get doc_id
|
|
62
|
+
temp_qa_df = sample_qa_df.copy(deep=True)
|
|
63
|
+
flatten_retrieval_gts = (
|
|
64
|
+
temp_qa_df["retrieval_gt"]
|
|
65
|
+
.apply(lambda x: list(itertools.chain.from_iterable(x)))
|
|
66
|
+
.tolist()
|
|
67
|
+
)
|
|
68
|
+
target_doc_ids = list(itertools.chain.from_iterable(flatten_retrieval_gts))
|
|
69
|
+
|
|
70
|
+
# make sample corpus data
|
|
71
|
+
sample_corpus_df = self.corpus_data.loc[
|
|
72
|
+
self.corpus_data["doc_id"].isin(target_doc_ids)
|
|
73
|
+
]
|
|
74
|
+
sample_corpus_df.reset_index(drop=True, inplace=True)
|
|
75
|
+
|
|
76
|
+
validate_qa_from_corpus_dataset(sample_qa_df, sample_corpus_df)
|
|
77
|
+
|
|
78
|
+
# start Evaluate at temp project directory
|
|
79
|
+
with (
|
|
80
|
+
tempfile.NamedTemporaryFile(suffix=".parquet", delete=False) as qa_path,
|
|
81
|
+
tempfile.NamedTemporaryFile(suffix=".parquet", delete=False) as corpus_path,
|
|
82
|
+
tempfile.TemporaryDirectory(ignore_cleanup_errors=True) as temp_project_dir,
|
|
83
|
+
):
|
|
84
|
+
sample_qa_df.to_parquet(qa_path.name, index=False)
|
|
85
|
+
sample_corpus_df.to_parquet(corpus_path.name, index=False)
|
|
86
|
+
|
|
87
|
+
evaluator = Evaluator(
|
|
88
|
+
qa_data_path=qa_path.name,
|
|
89
|
+
corpus_data_path=corpus_path.name,
|
|
90
|
+
project_dir=temp_project_dir,
|
|
91
|
+
)
|
|
92
|
+
evaluator.start_trial(yaml_path, skip_validation=True)
|
|
93
|
+
qa_path.close()
|
|
94
|
+
corpus_path.close()
|
|
95
|
+
os.unlink(qa_path.name)
|
|
96
|
+
os.unlink(corpus_path.name)
|
|
97
|
+
|
|
98
|
+
logger.info("Validation complete.")
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from typing import List
|
|
3
|
+
|
|
4
|
+
from autorag.support import dynamically_find_function
|
|
5
|
+
from autorag.utils.util import load_yaml_config
|
|
6
|
+
from autorag.vectordb.base import BaseVectorStore
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def get_support_vectordb(vectordb_name: str):
|
|
10
|
+
support_vectordb = {
|
|
11
|
+
"chroma": ("autorag.vectordb.chroma", "Chroma"),
|
|
12
|
+
"Chroma": ("autorag.vectordb.chroma", "Chroma"),
|
|
13
|
+
"milvus": ("autorag.vectordb.milvus", "Milvus"),
|
|
14
|
+
"Milvus": ("autorag.vectordb.milvus", "Milvus"),
|
|
15
|
+
"weaviate": ("autorag.vectordb.weaviate", "Weaviate"),
|
|
16
|
+
"Weaviate": ("autorag.vectordb.weaviate", "Weaviate"),
|
|
17
|
+
"pinecone": ("autorag.vectordb.pinecone", "Pinecone"),
|
|
18
|
+
"Pinecone": ("autorag.vectordb.pinecone", "Pinecone"),
|
|
19
|
+
"couchbase": ("autorag.vectordb.couchbase", "Couchbase"),
|
|
20
|
+
"Couchbase": ("autorag.vectordb.couchbase", "Couchbase"),
|
|
21
|
+
"qdrant": ("autorag.vectordb.qdrant", "Qdrant"),
|
|
22
|
+
"Qdrant": ("autorag.vectordb.qdrant", "Qdrant"),
|
|
23
|
+
}
|
|
24
|
+
return dynamically_find_function(vectordb_name, support_vectordb)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def load_vectordb(vectordb_name: str, **kwargs):
|
|
28
|
+
vectordb = get_support_vectordb(vectordb_name)
|
|
29
|
+
return vectordb(**kwargs)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def load_vectordb_from_yaml(yaml_path: str, vectordb_name: str, project_dir: str):
|
|
33
|
+
config_dict = load_yaml_config(yaml_path)
|
|
34
|
+
vectordb_list = config_dict.get("vectordb", [])
|
|
35
|
+
if len(vectordb_list) == 0 or vectordb_name == "default":
|
|
36
|
+
chroma_path = os.path.join(project_dir, "resources", "chroma")
|
|
37
|
+
return load_vectordb(
|
|
38
|
+
"chroma",
|
|
39
|
+
client_type="persistent",
|
|
40
|
+
embedding_model="openai",
|
|
41
|
+
collection_name="openai",
|
|
42
|
+
path=chroma_path,
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
target_dict = list(filter(lambda x: x["name"] == vectordb_name, vectordb_list))
|
|
46
|
+
target_dict[0].pop("name") # delete a name key
|
|
47
|
+
target_vectordb_name = target_dict[0].pop("db_type")
|
|
48
|
+
target_vectordb_params = target_dict[0]
|
|
49
|
+
return load_vectordb(target_vectordb_name, **target_vectordb_params)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def load_all_vectordb_from_yaml(
|
|
53
|
+
yaml_path: str, project_dir: str
|
|
54
|
+
) -> List[BaseVectorStore]:
|
|
55
|
+
config_dict = load_yaml_config(yaml_path)
|
|
56
|
+
vectordb_list = config_dict.get("vectordb", [])
|
|
57
|
+
if len(vectordb_list) == 0:
|
|
58
|
+
chroma_path = os.path.join(project_dir, "resources", "chroma")
|
|
59
|
+
return [
|
|
60
|
+
load_vectordb(
|
|
61
|
+
"chroma",
|
|
62
|
+
client_type="persistent",
|
|
63
|
+
embedding_model="openai",
|
|
64
|
+
collection_name="openai",
|
|
65
|
+
path=chroma_path,
|
|
66
|
+
)
|
|
67
|
+
]
|
|
68
|
+
|
|
69
|
+
result_vectordbs = []
|
|
70
|
+
for vectordb_dict in vectordb_list:
|
|
71
|
+
_ = vectordb_dict.pop("name")
|
|
72
|
+
vectordb_type = vectordb_dict.pop("db_type")
|
|
73
|
+
vectordb = load_vectordb(vectordb_type, **vectordb_dict)
|
|
74
|
+
result_vectordbs.append(vectordb)
|
|
75
|
+
return result_vectordbs
|
autorag/vectordb/base.py
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
from abc import abstractmethod
|
|
2
|
+
from typing import List, Tuple, Union
|
|
3
|
+
|
|
4
|
+
from llama_index.embeddings.openai import OpenAIEmbedding
|
|
5
|
+
|
|
6
|
+
from autorag.utils.util import openai_truncate_by_token
|
|
7
|
+
from autorag.embedding.base import EmbeddingModel
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class BaseVectorStore:
|
|
11
|
+
support_similarity_metrics = ["l2", "ip", "cosine"]
|
|
12
|
+
|
|
13
|
+
def __init__(
|
|
14
|
+
self,
|
|
15
|
+
embedding_model: Union[str, List[dict]],
|
|
16
|
+
similarity_metric: str = "cosine",
|
|
17
|
+
embedding_batch: int = 100,
|
|
18
|
+
):
|
|
19
|
+
self.embedding = EmbeddingModel.load(embedding_model)()
|
|
20
|
+
self.embedding_batch = embedding_batch
|
|
21
|
+
self.embedding.embed_batch_size = embedding_batch
|
|
22
|
+
assert similarity_metric in self.support_similarity_metrics, (
|
|
23
|
+
f"search method {similarity_metric} is not supported"
|
|
24
|
+
)
|
|
25
|
+
self.similarity_metric = similarity_metric
|
|
26
|
+
|
|
27
|
+
@abstractmethod
|
|
28
|
+
async def add(
|
|
29
|
+
self,
|
|
30
|
+
ids: List[str],
|
|
31
|
+
texts: List[str],
|
|
32
|
+
):
|
|
33
|
+
pass
|
|
34
|
+
|
|
35
|
+
@abstractmethod
|
|
36
|
+
def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
|
|
37
|
+
"""
|
|
38
|
+
Add the embeddings to the Vector DB.
|
|
39
|
+
"""
|
|
40
|
+
pass
|
|
41
|
+
|
|
42
|
+
@abstractmethod
|
|
43
|
+
async def query(
|
|
44
|
+
self, queries: List[str], top_k: int, **kwargs
|
|
45
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
46
|
+
pass
|
|
47
|
+
|
|
48
|
+
@abstractmethod
|
|
49
|
+
async def fetch(self, ids: List[str]) -> List[List[float]]:
|
|
50
|
+
"""
|
|
51
|
+
Fetch the embeddings of the ids.
|
|
52
|
+
"""
|
|
53
|
+
pass
|
|
54
|
+
|
|
55
|
+
@abstractmethod
|
|
56
|
+
async def is_exist(self, ids: List[str]) -> List[bool]:
|
|
57
|
+
"""
|
|
58
|
+
Check if the ids exist in the Vector DB.
|
|
59
|
+
"""
|
|
60
|
+
pass
|
|
61
|
+
|
|
62
|
+
@abstractmethod
|
|
63
|
+
async def delete(self, ids: List[str]):
|
|
64
|
+
pass
|
|
65
|
+
|
|
66
|
+
def truncated_inputs(self, inputs: List[str]) -> List[str]:
|
|
67
|
+
if isinstance(self.embedding, OpenAIEmbedding):
|
|
68
|
+
openai_embedding_limit = 8000
|
|
69
|
+
results = openai_truncate_by_token(
|
|
70
|
+
inputs, openai_embedding_limit, self.embedding.model_name
|
|
71
|
+
)
|
|
72
|
+
return results
|
|
73
|
+
return inputs
|
|
@@ -0,0 +1,118 @@
|
|
|
1
|
+
from typing import List, Optional, Dict, Tuple, Union
|
|
2
|
+
|
|
3
|
+
from chromadb import (
|
|
4
|
+
EphemeralClient,
|
|
5
|
+
PersistentClient,
|
|
6
|
+
DEFAULT_TENANT,
|
|
7
|
+
DEFAULT_DATABASE,
|
|
8
|
+
CloudClient,
|
|
9
|
+
AsyncHttpClient,
|
|
10
|
+
)
|
|
11
|
+
from chromadb.api.models.AsyncCollection import AsyncCollection
|
|
12
|
+
from chromadb.api.types import QueryResult
|
|
13
|
+
|
|
14
|
+
from autorag.utils.util import apply_recursive
|
|
15
|
+
from autorag.vectordb.base import BaseVectorStore
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class Chroma(BaseVectorStore):
|
|
19
|
+
def __init__(
|
|
20
|
+
self,
|
|
21
|
+
embedding_model: Union[str, List[dict]],
|
|
22
|
+
collection_name: str,
|
|
23
|
+
embedding_batch: int = 100,
|
|
24
|
+
client_type: str = "persistent",
|
|
25
|
+
similarity_metric: str = "cosine",
|
|
26
|
+
path: str = None,
|
|
27
|
+
host: str = "localhost",
|
|
28
|
+
port: int = 8000,
|
|
29
|
+
ssl: bool = False,
|
|
30
|
+
headers: Optional[Dict[str, str]] = None,
|
|
31
|
+
api_key: Optional[str] = None,
|
|
32
|
+
tenant: str = DEFAULT_TENANT,
|
|
33
|
+
database: str = DEFAULT_DATABASE,
|
|
34
|
+
):
|
|
35
|
+
super().__init__(embedding_model, similarity_metric, embedding_batch)
|
|
36
|
+
if client_type == "ephemeral":
|
|
37
|
+
self.client = EphemeralClient(tenant=tenant, database=database)
|
|
38
|
+
elif client_type == "persistent":
|
|
39
|
+
assert path is not None, "path must be provided for persistent client"
|
|
40
|
+
self.client = PersistentClient(path=path, tenant=tenant, database=database)
|
|
41
|
+
elif client_type == "http":
|
|
42
|
+
self.client = AsyncHttpClient(
|
|
43
|
+
host=host,
|
|
44
|
+
port=port,
|
|
45
|
+
ssl=ssl,
|
|
46
|
+
headers=headers,
|
|
47
|
+
tenant=tenant,
|
|
48
|
+
database=database,
|
|
49
|
+
)
|
|
50
|
+
elif client_type == "cloud":
|
|
51
|
+
self.client = CloudClient(
|
|
52
|
+
tenant=tenant,
|
|
53
|
+
database=database,
|
|
54
|
+
api_key=api_key,
|
|
55
|
+
)
|
|
56
|
+
else:
|
|
57
|
+
raise ValueError(
|
|
58
|
+
f"client_type {client_type} is not supported\n"
|
|
59
|
+
"supported client types are: ephemeral, persistent, http, cloud"
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
self.collection = self.client.get_or_create_collection(
|
|
63
|
+
name=collection_name,
|
|
64
|
+
metadata={"hnsw:space": similarity_metric},
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
async def add(self, ids: List[str], texts: List[str]):
|
|
68
|
+
texts = self.truncated_inputs(texts)
|
|
69
|
+
text_embeddings = await self.embedding.aget_text_embedding_batch(texts)
|
|
70
|
+
if isinstance(self.collection, AsyncCollection):
|
|
71
|
+
await self.collection.add(ids=ids, embeddings=text_embeddings)
|
|
72
|
+
else:
|
|
73
|
+
self.collection.add(ids=ids, embeddings=text_embeddings)
|
|
74
|
+
|
|
75
|
+
def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
|
|
76
|
+
self.collection.add(ids=ids, embeddings=embeddings)
|
|
77
|
+
|
|
78
|
+
async def fetch(self, ids: List[str]) -> List[List[float]]:
|
|
79
|
+
if isinstance(self.collection, AsyncCollection):
|
|
80
|
+
fetch_result = await self.collection.get(ids, include=["embeddings"])
|
|
81
|
+
else:
|
|
82
|
+
fetch_result = self.collection.get(ids, include=["embeddings"])
|
|
83
|
+
fetch_embeddings = fetch_result["embeddings"]
|
|
84
|
+
return fetch_embeddings
|
|
85
|
+
|
|
86
|
+
async def is_exist(self, ids: List[str]) -> List[bool]:
|
|
87
|
+
if isinstance(self.collection, AsyncCollection):
|
|
88
|
+
fetched_result = await self.collection.get(ids, include=[])
|
|
89
|
+
else:
|
|
90
|
+
fetched_result = self.collection.get(ids, include=[])
|
|
91
|
+
existed_ids = fetched_result["ids"]
|
|
92
|
+
return list(map(lambda x: x in existed_ids, ids))
|
|
93
|
+
|
|
94
|
+
async def query(
|
|
95
|
+
self, queries: List[str], top_k: int, **kwargs
|
|
96
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
97
|
+
queries = self.truncated_inputs(queries)
|
|
98
|
+
query_embeddings: List[
|
|
99
|
+
List[float]
|
|
100
|
+
] = await self.embedding.aget_text_embedding_batch(queries)
|
|
101
|
+
if isinstance(self.collection, AsyncCollection):
|
|
102
|
+
query_result: QueryResult = await self.collection.query(
|
|
103
|
+
query_embeddings=query_embeddings, n_results=top_k
|
|
104
|
+
)
|
|
105
|
+
else:
|
|
106
|
+
query_result: QueryResult = self.collection.query(
|
|
107
|
+
query_embeddings=query_embeddings, n_results=top_k
|
|
108
|
+
)
|
|
109
|
+
ids = query_result["ids"]
|
|
110
|
+
scores = query_result["distances"]
|
|
111
|
+
scores = apply_recursive(lambda x: 1 - x, scores)
|
|
112
|
+
return ids, scores
|
|
113
|
+
|
|
114
|
+
async def delete(self, ids: List[str]):
|
|
115
|
+
if isinstance(self.collection, AsyncCollection):
|
|
116
|
+
await self.collection.delete(ids)
|
|
117
|
+
else:
|
|
118
|
+
self.collection.delete(ids)
|
|
@@ -0,0 +1,239 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
from datetime import timedelta
|
|
4
|
+
|
|
5
|
+
from couchbase.auth import PasswordAuthenticator
|
|
6
|
+
from couchbase.cluster import Cluster
|
|
7
|
+
from couchbase.options import ClusterOptions
|
|
8
|
+
|
|
9
|
+
from typing import List, Tuple, Optional, Union
|
|
10
|
+
|
|
11
|
+
from autorag.utils.util import make_batch
|
|
12
|
+
from autorag.vectordb import BaseVectorStore
|
|
13
|
+
|
|
14
|
+
logger = logging.getLogger("AutoRAG")
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Couchbase(BaseVectorStore):
|
|
18
|
+
def __init__(
|
|
19
|
+
self,
|
|
20
|
+
embedding_model: Union[str, List[dict]],
|
|
21
|
+
bucket_name: str,
|
|
22
|
+
scope_name: str,
|
|
23
|
+
collection_name: str,
|
|
24
|
+
index_name: str,
|
|
25
|
+
embedding_batch: int = 100,
|
|
26
|
+
connection_string: str = "",
|
|
27
|
+
username: str = "",
|
|
28
|
+
password: str = "",
|
|
29
|
+
ingest_batch: int = 100,
|
|
30
|
+
text_key: Optional[str] = "text",
|
|
31
|
+
embedding_key: Optional[str] = "embedding",
|
|
32
|
+
scoped_index: bool = True,
|
|
33
|
+
):
|
|
34
|
+
super().__init__(
|
|
35
|
+
embedding_model=embedding_model,
|
|
36
|
+
similarity_metric="ip",
|
|
37
|
+
embedding_batch=embedding_batch,
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
self.index_name = index_name
|
|
41
|
+
self.bucket_name = bucket_name
|
|
42
|
+
self.scope_name = scope_name
|
|
43
|
+
self.collection_name = collection_name
|
|
44
|
+
self.scoped_index = scoped_index
|
|
45
|
+
self.text_key = text_key
|
|
46
|
+
self.embedding_key = embedding_key
|
|
47
|
+
self.ingest_batch = ingest_batch
|
|
48
|
+
|
|
49
|
+
auth = PasswordAuthenticator(username, password)
|
|
50
|
+
self.cluster = Cluster(connection_string, ClusterOptions(auth))
|
|
51
|
+
|
|
52
|
+
# Wait until the cluster is ready for use.
|
|
53
|
+
self.cluster.wait_until_ready(timedelta(seconds=5))
|
|
54
|
+
|
|
55
|
+
# Check if the bucket exists
|
|
56
|
+
if not self._check_bucket_exists():
|
|
57
|
+
raise ValueError(
|
|
58
|
+
f"Bucket {self.bucket_name} does not exist. "
|
|
59
|
+
" Please create the bucket before searching."
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
try:
|
|
63
|
+
self.bucket = self.cluster.bucket(self.bucket_name)
|
|
64
|
+
self.scope = self.bucket.scope(self.scope_name)
|
|
65
|
+
self.collection = self.scope.collection(self.collection_name)
|
|
66
|
+
except Exception as e:
|
|
67
|
+
raise ValueError(
|
|
68
|
+
"Error connecting to couchbase. "
|
|
69
|
+
"Please check the connection and credentials."
|
|
70
|
+
) from e
|
|
71
|
+
|
|
72
|
+
# Check if the index exists. Throws ValueError if it doesn't
|
|
73
|
+
try:
|
|
74
|
+
self._check_index_exists()
|
|
75
|
+
except Exception:
|
|
76
|
+
raise
|
|
77
|
+
|
|
78
|
+
# Reinitialize to ensure a consistent state
|
|
79
|
+
self.bucket = self.cluster.bucket(self.bucket_name)
|
|
80
|
+
self.scope = self.bucket.scope(self.scope_name)
|
|
81
|
+
self.collection = self.scope.collection(self.collection_name)
|
|
82
|
+
|
|
83
|
+
async def add(self, ids: List[str], texts: List[str]):
|
|
84
|
+
from couchbase.exceptions import DocumentExistsException
|
|
85
|
+
|
|
86
|
+
texts = self.truncated_inputs(texts)
|
|
87
|
+
text_embeddings: List[
|
|
88
|
+
List[float]
|
|
89
|
+
] = await self.embedding.aget_text_embedding_batch(texts)
|
|
90
|
+
|
|
91
|
+
documents_to_insert = []
|
|
92
|
+
for _id, text, embedding in zip(ids, texts, text_embeddings):
|
|
93
|
+
doc = {
|
|
94
|
+
self.text_key: text,
|
|
95
|
+
self.embedding_key: embedding,
|
|
96
|
+
}
|
|
97
|
+
documents_to_insert.append({_id: doc})
|
|
98
|
+
|
|
99
|
+
batch_documents_to_insert = make_batch(documents_to_insert, self.ingest_batch)
|
|
100
|
+
|
|
101
|
+
for batch in batch_documents_to_insert:
|
|
102
|
+
insert_batch = {}
|
|
103
|
+
for doc in batch:
|
|
104
|
+
insert_batch.update(doc)
|
|
105
|
+
try:
|
|
106
|
+
self.collection.upsert_multi(insert_batch)
|
|
107
|
+
except DocumentExistsException as e:
|
|
108
|
+
logger.debug(f"Document already exists: {e}")
|
|
109
|
+
|
|
110
|
+
def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
|
|
111
|
+
from couchbase.exceptions import DocumentExistsException
|
|
112
|
+
|
|
113
|
+
documents_to_insert = []
|
|
114
|
+
for _id, embedding in zip(ids, embeddings):
|
|
115
|
+
doc = {
|
|
116
|
+
self.embedding_key: embedding,
|
|
117
|
+
}
|
|
118
|
+
documents_to_insert.append({_id: doc})
|
|
119
|
+
|
|
120
|
+
batch_documents_to_insert = make_batch(documents_to_insert, self.ingest_batch)
|
|
121
|
+
|
|
122
|
+
for batch in batch_documents_to_insert:
|
|
123
|
+
insert_batch = {}
|
|
124
|
+
for doc in batch:
|
|
125
|
+
insert_batch.update(doc)
|
|
126
|
+
try:
|
|
127
|
+
self.collection.upsert_multi(insert_batch)
|
|
128
|
+
except DocumentExistsException as e:
|
|
129
|
+
logger.debug(f"Document already exists: {e}")
|
|
130
|
+
|
|
131
|
+
async def fetch(self, ids: List[str]) -> List[List[float]]:
|
|
132
|
+
# Fetch vectors by IDs
|
|
133
|
+
fetched_result = self.collection.get_multi(ids)
|
|
134
|
+
fetched_vectors = {
|
|
135
|
+
k: v.value[f"{self.embedding_key}"]
|
|
136
|
+
for k, v in fetched_result.results.items()
|
|
137
|
+
}
|
|
138
|
+
return list(map(lambda x: fetched_vectors[x], ids))
|
|
139
|
+
|
|
140
|
+
async def is_exist(self, ids: List[str]) -> List[bool]:
|
|
141
|
+
existed_result = self.collection.exists_multi(ids)
|
|
142
|
+
existed_ids = {k: v.exists for k, v in existed_result.results.items()}
|
|
143
|
+
return list(map(lambda x: existed_ids[x], ids))
|
|
144
|
+
|
|
145
|
+
async def query(
|
|
146
|
+
self, queries: List[str], top_k: int, **kwargs
|
|
147
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
148
|
+
import couchbase.search as search
|
|
149
|
+
from couchbase.options import SearchOptions
|
|
150
|
+
from couchbase.vector_search import VectorQuery, VectorSearch
|
|
151
|
+
|
|
152
|
+
queries = self.truncated_inputs(queries)
|
|
153
|
+
query_embeddings: List[
|
|
154
|
+
List[float]
|
|
155
|
+
] = await self.embedding.aget_text_embedding_batch(queries)
|
|
156
|
+
|
|
157
|
+
ids, scores = [], []
|
|
158
|
+
for query_embedding in query_embeddings:
|
|
159
|
+
# Create Search Request
|
|
160
|
+
search_req = search.SearchRequest.create(
|
|
161
|
+
VectorSearch.from_vector_query(
|
|
162
|
+
VectorQuery(
|
|
163
|
+
self.embedding_key,
|
|
164
|
+
query_embedding,
|
|
165
|
+
top_k,
|
|
166
|
+
)
|
|
167
|
+
)
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
# Search
|
|
171
|
+
if self.scoped_index:
|
|
172
|
+
search_iter = self.scope.search(
|
|
173
|
+
self.index_name,
|
|
174
|
+
search_req,
|
|
175
|
+
SearchOptions(limit=top_k),
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
else:
|
|
179
|
+
search_iter = self.cluster.search(
|
|
180
|
+
self.index_name,
|
|
181
|
+
search_req,
|
|
182
|
+
SearchOptions(limit=top_k),
|
|
183
|
+
)
|
|
184
|
+
|
|
185
|
+
# Parse the search results
|
|
186
|
+
# search_iter.rows() can only be iterated once.
|
|
187
|
+
id_list, score_list = [], []
|
|
188
|
+
for result in search_iter.rows():
|
|
189
|
+
id_list.append(result.id)
|
|
190
|
+
score_list.append(result.score)
|
|
191
|
+
|
|
192
|
+
ids.append(id_list)
|
|
193
|
+
scores.append(score_list)
|
|
194
|
+
|
|
195
|
+
return ids, scores
|
|
196
|
+
|
|
197
|
+
async def delete(self, ids: List[str]):
|
|
198
|
+
self.collection.remove_multi(ids)
|
|
199
|
+
|
|
200
|
+
def _check_bucket_exists(self) -> bool:
|
|
201
|
+
"""Check if the bucket exists in the linked Couchbase cluster.
|
|
202
|
+
|
|
203
|
+
Returns:
|
|
204
|
+
True if the bucket exists
|
|
205
|
+
"""
|
|
206
|
+
bucket_manager = self.cluster.buckets()
|
|
207
|
+
try:
|
|
208
|
+
bucket_manager.get_bucket(self.bucket_name)
|
|
209
|
+
return True
|
|
210
|
+
except Exception as e:
|
|
211
|
+
logger.debug("Error checking if bucket exists:", e)
|
|
212
|
+
return False
|
|
213
|
+
|
|
214
|
+
def _check_index_exists(self) -> bool:
|
|
215
|
+
"""Check if the Search index exists in the linked Couchbase cluster
|
|
216
|
+
Returns:
|
|
217
|
+
bool: True if the index exists, False otherwise.
|
|
218
|
+
Raises a ValueError if the index does not exist.
|
|
219
|
+
"""
|
|
220
|
+
if self.scoped_index:
|
|
221
|
+
all_indexes = [
|
|
222
|
+
index.name for index in self.scope.search_indexes().get_all_indexes()
|
|
223
|
+
]
|
|
224
|
+
if self.index_name not in all_indexes:
|
|
225
|
+
raise ValueError(
|
|
226
|
+
f"Index {self.index_name} does not exist. "
|
|
227
|
+
" Please create the index before searching."
|
|
228
|
+
)
|
|
229
|
+
else:
|
|
230
|
+
all_indexes = [
|
|
231
|
+
index.name for index in self.cluster.search_indexes().get_all_indexes()
|
|
232
|
+
]
|
|
233
|
+
if self.index_name not in all_indexes:
|
|
234
|
+
raise ValueError(
|
|
235
|
+
f"Index {self.index_name} does not exist. "
|
|
236
|
+
" Please create the index before searching."
|
|
237
|
+
)
|
|
238
|
+
|
|
239
|
+
return True
|