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,339 @@
|
|
|
1
|
+
import itertools
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
from typing import List, Tuple, Optional
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import pandas as pd
|
|
8
|
+
from llama_index.core.embeddings import BaseEmbedding
|
|
9
|
+
from llama_index.embeddings.openai import OpenAIEmbedding
|
|
10
|
+
|
|
11
|
+
from autorag.evaluation.metric.util import (
|
|
12
|
+
calculate_l2_distance,
|
|
13
|
+
calculate_inner_product,
|
|
14
|
+
calculate_cosine_similarity,
|
|
15
|
+
)
|
|
16
|
+
from autorag.nodes.retrieval.base import evenly_distribute_passages, BaseRetrieval
|
|
17
|
+
from autorag.utils import (
|
|
18
|
+
validate_corpus_dataset,
|
|
19
|
+
cast_corpus_dataset,
|
|
20
|
+
cast_qa_dataset,
|
|
21
|
+
validate_qa_dataset,
|
|
22
|
+
)
|
|
23
|
+
from autorag.utils.util import (
|
|
24
|
+
get_event_loop,
|
|
25
|
+
process_batch,
|
|
26
|
+
openai_truncate_by_token,
|
|
27
|
+
flatten_apply,
|
|
28
|
+
result_to_dataframe,
|
|
29
|
+
pop_params,
|
|
30
|
+
fetch_contents,
|
|
31
|
+
empty_cuda_cache,
|
|
32
|
+
convert_inputs_to_list,
|
|
33
|
+
make_batch,
|
|
34
|
+
)
|
|
35
|
+
from autorag.vectordb import load_vectordb_from_yaml
|
|
36
|
+
from autorag.vectordb.base import BaseVectorStore
|
|
37
|
+
|
|
38
|
+
logger = logging.getLogger("AutoRAG")
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class VectorDB(BaseRetrieval):
|
|
42
|
+
def __init__(self, project_dir: str, vectordb: str = "default", **kwargs):
|
|
43
|
+
"""
|
|
44
|
+
Initialize VectorDB retrieval node.
|
|
45
|
+
|
|
46
|
+
:param project_dir: The project directory path.
|
|
47
|
+
:param vectordb: The vectordb name.
|
|
48
|
+
You must configure the vectordb name in the config.yaml file.
|
|
49
|
+
If you don't configure, it uses the default vectordb.
|
|
50
|
+
:param kwargs: The optional arguments.
|
|
51
|
+
Not affected in the init method.
|
|
52
|
+
"""
|
|
53
|
+
super().__init__(project_dir)
|
|
54
|
+
|
|
55
|
+
vectordb_config_path = os.path.join(self.resources_dir, "vectordb.yaml")
|
|
56
|
+
self.vector_store = load_vectordb_from_yaml(
|
|
57
|
+
vectordb_config_path, vectordb, project_dir
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
self.embedding_model = self.vector_store.embedding
|
|
61
|
+
|
|
62
|
+
def __del__(self):
|
|
63
|
+
del self.vector_store
|
|
64
|
+
del self.embedding_model
|
|
65
|
+
empty_cuda_cache()
|
|
66
|
+
super().__del__()
|
|
67
|
+
|
|
68
|
+
@result_to_dataframe(
|
|
69
|
+
[
|
|
70
|
+
"retrieved_contents_semantic",
|
|
71
|
+
"retrieved_ids_semantic",
|
|
72
|
+
"retrieve_scores_semantic",
|
|
73
|
+
]
|
|
74
|
+
)
|
|
75
|
+
def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
|
|
76
|
+
queries = self.cast_to_run(previous_result)
|
|
77
|
+
pure_params = pop_params(self._pure, kwargs)
|
|
78
|
+
ids, scores = self._pure(queries, **pure_params)
|
|
79
|
+
contents = fetch_contents(self.corpus_df, ids)
|
|
80
|
+
return contents, ids, scores
|
|
81
|
+
|
|
82
|
+
def _pure(
|
|
83
|
+
self,
|
|
84
|
+
queries: List[List[str]],
|
|
85
|
+
top_k: int,
|
|
86
|
+
embedding_batch: int = 128,
|
|
87
|
+
ids: Optional[List[List[str]]] = None,
|
|
88
|
+
) -> Tuple[List[List[str]], List[List[float]]]:
|
|
89
|
+
"""
|
|
90
|
+
VectorDB retrieval function.
|
|
91
|
+
You have to get a chroma collection that is already ingested.
|
|
92
|
+
You have to get an embedding model that is already used in ingesting.
|
|
93
|
+
|
|
94
|
+
:param queries: 2-d list of query strings.
|
|
95
|
+
Each element of the list is a query strings of each row.
|
|
96
|
+
:param top_k: The number of passages to be retrieved.
|
|
97
|
+
:param embedding_batch: The number of queries to be processed in parallel.
|
|
98
|
+
This is used to prevent API error at the query embedding.
|
|
99
|
+
Default is 128.
|
|
100
|
+
:param ids: The optional list of ids that you want to retrieve.
|
|
101
|
+
You don't need to specify this in the general use cases.
|
|
102
|
+
Default is None.
|
|
103
|
+
|
|
104
|
+
:return: The 2-d list contains a list of passage ids that retrieved from vectordb and 2-d list of its scores.
|
|
105
|
+
It will be a length of queries. And each element has a length of top_k.
|
|
106
|
+
"""
|
|
107
|
+
# if ids are specified, fetch the ids score from Chroma
|
|
108
|
+
if ids is not None:
|
|
109
|
+
return self.__get_ids_scores(queries, ids, embedding_batch)
|
|
110
|
+
|
|
111
|
+
# run async vector_db_pure function
|
|
112
|
+
tasks = [
|
|
113
|
+
vectordb_pure(query_list, top_k, self.vector_store)
|
|
114
|
+
for query_list in queries
|
|
115
|
+
]
|
|
116
|
+
loop = get_event_loop()
|
|
117
|
+
results = loop.run_until_complete(
|
|
118
|
+
process_batch(tasks, batch_size=embedding_batch)
|
|
119
|
+
)
|
|
120
|
+
id_result = list(map(lambda x: x[0], results))
|
|
121
|
+
score_result = list(map(lambda x: x[1], results))
|
|
122
|
+
return id_result, score_result
|
|
123
|
+
|
|
124
|
+
def __get_ids_scores(self, queries, ids, embedding_batch: int):
|
|
125
|
+
# truncate queries and embedding execution here.
|
|
126
|
+
openai_embedding_limit = 8000
|
|
127
|
+
if isinstance(self.embedding_model, OpenAIEmbedding):
|
|
128
|
+
queries = list(
|
|
129
|
+
map(
|
|
130
|
+
lambda query_list: openai_truncate_by_token(
|
|
131
|
+
query_list,
|
|
132
|
+
openai_embedding_limit,
|
|
133
|
+
self.embedding_model.model_name,
|
|
134
|
+
),
|
|
135
|
+
queries,
|
|
136
|
+
)
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
query_embeddings = flatten_apply(
|
|
140
|
+
run_query_embedding_batch,
|
|
141
|
+
queries,
|
|
142
|
+
embedding_model=self.embedding_model,
|
|
143
|
+
batch_size=embedding_batch,
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
loop = get_event_loop()
|
|
147
|
+
|
|
148
|
+
async def run_fetch(ids):
|
|
149
|
+
final_result = []
|
|
150
|
+
for id_list in ids:
|
|
151
|
+
if len(id_list) == 0:
|
|
152
|
+
final_result.append([])
|
|
153
|
+
else:
|
|
154
|
+
result = await self.vector_store.fetch(id_list)
|
|
155
|
+
final_result.append(result)
|
|
156
|
+
return final_result
|
|
157
|
+
|
|
158
|
+
content_embeddings = loop.run_until_complete(run_fetch(ids))
|
|
159
|
+
|
|
160
|
+
score_result = list(
|
|
161
|
+
map(
|
|
162
|
+
lambda query_embedding_list, content_embedding_list: get_id_scores(
|
|
163
|
+
query_embedding_list,
|
|
164
|
+
content_embedding_list,
|
|
165
|
+
similarity_metric=self.vector_store.similarity_metric,
|
|
166
|
+
),
|
|
167
|
+
query_embeddings,
|
|
168
|
+
content_embeddings,
|
|
169
|
+
)
|
|
170
|
+
)
|
|
171
|
+
return ids, score_result
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
async def vectordb_pure(
|
|
175
|
+
queries: List[str], top_k: int, vectordb: BaseVectorStore
|
|
176
|
+
) -> Tuple[List[str], List[float]]:
|
|
177
|
+
"""
|
|
178
|
+
Async VectorDB retrieval function.
|
|
179
|
+
Its usage is for async retrieval of vector_db row by row.
|
|
180
|
+
|
|
181
|
+
:param query_embeddings: A list of query embeddings.
|
|
182
|
+
:param top_k: The number of passages to be retrieved.
|
|
183
|
+
:param vectordb: The vector store instance.
|
|
184
|
+
:return: The tuple contains a list of passage ids that are retrieved from vectordb and a list of its scores.
|
|
185
|
+
"""
|
|
186
|
+
id_result, score_result = await vectordb.query(queries=queries, top_k=top_k)
|
|
187
|
+
|
|
188
|
+
# Distribute passages evenly
|
|
189
|
+
id_result, score_result = evenly_distribute_passages(id_result, score_result, top_k)
|
|
190
|
+
# sort id_result and score_result by score
|
|
191
|
+
result = [
|
|
192
|
+
(_id, score)
|
|
193
|
+
for score, _id in sorted(
|
|
194
|
+
zip(score_result, id_result), key=lambda pair: pair[0], reverse=True
|
|
195
|
+
)
|
|
196
|
+
]
|
|
197
|
+
id_result, score_result = zip(*result)
|
|
198
|
+
return list(id_result), list(score_result)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
async def filter_exist_ids(
|
|
202
|
+
vectordb: BaseVectorStore,
|
|
203
|
+
corpus_data: pd.DataFrame,
|
|
204
|
+
) -> pd.DataFrame:
|
|
205
|
+
corpus_data = cast_corpus_dataset(corpus_data)
|
|
206
|
+
validate_corpus_dataset(corpus_data)
|
|
207
|
+
ids = corpus_data["doc_id"].tolist()
|
|
208
|
+
|
|
209
|
+
# Query the collection to check if IDs already exist
|
|
210
|
+
existed_bool_list = await vectordb.is_exist(ids=ids)
|
|
211
|
+
# Assuming 'ids' is the key in the response
|
|
212
|
+
new_passage = corpus_data[~pd.Series(existed_bool_list)]
|
|
213
|
+
return new_passage
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
async def filter_exist_ids_from_retrieval_gt(
|
|
217
|
+
vectordb: BaseVectorStore,
|
|
218
|
+
qa_data: pd.DataFrame,
|
|
219
|
+
corpus_data: pd.DataFrame,
|
|
220
|
+
) -> pd.DataFrame:
|
|
221
|
+
qa_data = cast_qa_dataset(qa_data)
|
|
222
|
+
validate_qa_dataset(qa_data)
|
|
223
|
+
corpus_data = cast_corpus_dataset(corpus_data)
|
|
224
|
+
validate_corpus_dataset(corpus_data)
|
|
225
|
+
retrieval_gt = (
|
|
226
|
+
qa_data["retrieval_gt"]
|
|
227
|
+
.apply(lambda x: list(itertools.chain.from_iterable(x)))
|
|
228
|
+
.tolist()
|
|
229
|
+
)
|
|
230
|
+
retrieval_gt = list(itertools.chain.from_iterable(retrieval_gt))
|
|
231
|
+
retrieval_gt = list(set(retrieval_gt))
|
|
232
|
+
|
|
233
|
+
existed_bool_list = await vectordb.is_exist(ids=retrieval_gt)
|
|
234
|
+
add_ids = []
|
|
235
|
+
for ret_gt, is_exist in zip(retrieval_gt, existed_bool_list):
|
|
236
|
+
if not is_exist:
|
|
237
|
+
add_ids.append(ret_gt)
|
|
238
|
+
new_passage = corpus_data[corpus_data["doc_id"].isin(add_ids)]
|
|
239
|
+
return new_passage
|
|
240
|
+
|
|
241
|
+
|
|
242
|
+
async def vectordb_ingest_api(
|
|
243
|
+
vectordb: BaseVectorStore,
|
|
244
|
+
corpus_data: pd.DataFrame,
|
|
245
|
+
):
|
|
246
|
+
"""
|
|
247
|
+
Ingest given corpus data to the vectordb.
|
|
248
|
+
It truncates corpus content when the embedding model is OpenAIEmbedding to the 8000 tokens.
|
|
249
|
+
Plus, when the corpus content is empty (whitespace), it will be ignored.
|
|
250
|
+
And if there is a document id that already exists in the collection, it will be ignored.
|
|
251
|
+
|
|
252
|
+
:param vectordb: A vector stores instance that you want to ingest.
|
|
253
|
+
:param corpus_data: The corpus data that contains doc_id and contents columns.
|
|
254
|
+
"""
|
|
255
|
+
embedding_batch = vectordb.embedding_batch
|
|
256
|
+
if not corpus_data.empty:
|
|
257
|
+
new_contents = corpus_data["contents"].tolist()
|
|
258
|
+
new_ids = corpus_data["doc_id"].tolist()
|
|
259
|
+
content_batches = make_batch(new_contents, embedding_batch)
|
|
260
|
+
id_batches = make_batch(new_ids, embedding_batch)
|
|
261
|
+
for content_batch, id_batch in zip(content_batches, id_batches):
|
|
262
|
+
await vectordb.add(ids=id_batch, texts=content_batch)
|
|
263
|
+
|
|
264
|
+
|
|
265
|
+
def vectordb_ingest_huggingface(
|
|
266
|
+
vectordb: BaseVectorStore,
|
|
267
|
+
corpus_data: pd.DataFrame,
|
|
268
|
+
):
|
|
269
|
+
"""
|
|
270
|
+
Ingest given corpus data to the vectordb using local model.
|
|
271
|
+
When the corpus content is empty (whitespace), it will be ignored.
|
|
272
|
+
And if there is a document id that already exists in the collection, it will be ignored.
|
|
273
|
+
|
|
274
|
+
:param vectordb: A vector stores instance that you want to ingest.
|
|
275
|
+
:param corpus_data: The corpus data that contains doc_id and contents columns.
|
|
276
|
+
"""
|
|
277
|
+
embedding_batch_size = vectordb.embedding_batch
|
|
278
|
+
embedding_model = vectordb.embedding._model
|
|
279
|
+
if corpus_data.empty:
|
|
280
|
+
logger.warning("The corpus data is empty. Nothing to ingest.")
|
|
281
|
+
return
|
|
282
|
+
new_contents = corpus_data["contents"].tolist()
|
|
283
|
+
new_ids = corpus_data["doc_id"].tolist()
|
|
284
|
+
logger.info("Start embedding corpus data with huggingface model.")
|
|
285
|
+
embeddings = embedding_model.encode(
|
|
286
|
+
new_contents,
|
|
287
|
+
batch_size=embedding_batch_size,
|
|
288
|
+
normalize_embeddings=vectordb.embedding.normalize,
|
|
289
|
+
show_progress_bar=True,
|
|
290
|
+
)
|
|
291
|
+
vectordb.add_embedding(new_ids, embeddings)
|
|
292
|
+
logger.info("Finish embedding & ingesting corpus data with huggingface model.")
|
|
293
|
+
|
|
294
|
+
|
|
295
|
+
def run_query_embedding_batch(
|
|
296
|
+
queries: List[str], embedding_model: BaseEmbedding, batch_size: int
|
|
297
|
+
) -> List[List[float]]:
|
|
298
|
+
result = []
|
|
299
|
+
for i in range(0, len(queries), batch_size):
|
|
300
|
+
batch = queries[i : i + batch_size]
|
|
301
|
+
embeddings = embedding_model.get_text_embedding_batch(batch)
|
|
302
|
+
result.extend(embeddings)
|
|
303
|
+
return result
|
|
304
|
+
|
|
305
|
+
|
|
306
|
+
@convert_inputs_to_list
|
|
307
|
+
def get_id_scores( # To find the uncalculated score when fuse the scores for the hybrid retrieval
|
|
308
|
+
query_embeddings: List[
|
|
309
|
+
List[float]
|
|
310
|
+
], # `queries` is input. This is one user input query.
|
|
311
|
+
content_embeddings: List[List[float]],
|
|
312
|
+
similarity_metric: str,
|
|
313
|
+
) -> List[
|
|
314
|
+
float
|
|
315
|
+
]: # The most high scores among each query. The length of a result is the same as the contents length.
|
|
316
|
+
"""
|
|
317
|
+
Calculate the highest similarity scores between query embeddings and content embeddings.
|
|
318
|
+
|
|
319
|
+
:param query_embeddings: A list of lists containing query embeddings.
|
|
320
|
+
:param content_embeddings: A list of lists containing content embeddings.
|
|
321
|
+
:param similarity_metric: The similarity metric to use ('l2', 'ip', or 'cosine').
|
|
322
|
+
:return: A list of the highest similarity scores for each content embedding.
|
|
323
|
+
"""
|
|
324
|
+
metric_func_dict = {
|
|
325
|
+
"l2": lambda x, y: 1 - calculate_l2_distance(x, y),
|
|
326
|
+
"ip": calculate_inner_product,
|
|
327
|
+
"cosine": calculate_cosine_similarity,
|
|
328
|
+
}
|
|
329
|
+
metric_func = metric_func_dict[similarity_metric]
|
|
330
|
+
|
|
331
|
+
result = []
|
|
332
|
+
for content_embedding in content_embeddings:
|
|
333
|
+
scores = []
|
|
334
|
+
for query_embedding in query_embeddings:
|
|
335
|
+
scores.append(
|
|
336
|
+
metric_func(np.array(query_embedding), np.array(content_embedding))
|
|
337
|
+
)
|
|
338
|
+
result.append(max(scores))
|
|
339
|
+
return result
|
autorag/nodes/util.py
ADDED
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
from typing import Optional, Dict
|
|
2
|
+
|
|
3
|
+
from autorag.support import get_support_modules
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def make_generator_callable_param(generator_dict: Optional[Dict]):
|
|
7
|
+
if "generator_module_type" not in generator_dict.keys():
|
|
8
|
+
generator_dict = {
|
|
9
|
+
"generator_module_type": "llama_index_llm",
|
|
10
|
+
"llm": "openai",
|
|
11
|
+
"model": "gpt-4o-mini",
|
|
12
|
+
}
|
|
13
|
+
module_str = generator_dict.pop("generator_module_type")
|
|
14
|
+
module_class = get_support_modules(module_str)
|
|
15
|
+
module_param = generator_dict
|
|
16
|
+
return module_class, module_param
|
autorag/parser.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import os
|
|
3
|
+
import shutil
|
|
4
|
+
from typing import Optional
|
|
5
|
+
|
|
6
|
+
from autorag.data.parse.run import run_parser
|
|
7
|
+
from autorag.data.utils.util import load_yaml, get_param_combinations
|
|
8
|
+
|
|
9
|
+
logger = logging.getLogger("AutoRAG")
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class Parser:
|
|
13
|
+
def __init__(self, data_path_glob: str, project_dir: Optional[str] = None):
|
|
14
|
+
self.data_path_glob = data_path_glob
|
|
15
|
+
self.project_dir = project_dir if project_dir is not None else os.getcwd()
|
|
16
|
+
|
|
17
|
+
def start_parsing(self, yaml_path: str, all_files: bool = False):
|
|
18
|
+
if not os.path.exists(self.project_dir):
|
|
19
|
+
os.makedirs(self.project_dir)
|
|
20
|
+
|
|
21
|
+
# copy yaml file to project directory
|
|
22
|
+
shutil.copy(yaml_path, os.path.join(self.project_dir, "parse_config.yaml"))
|
|
23
|
+
|
|
24
|
+
# load yaml file
|
|
25
|
+
modules = load_yaml(yaml_path)
|
|
26
|
+
|
|
27
|
+
input_modules, input_params = get_param_combinations(modules)
|
|
28
|
+
|
|
29
|
+
logger.info("Parsing Start...")
|
|
30
|
+
run_parser(
|
|
31
|
+
modules=input_modules,
|
|
32
|
+
module_params=input_params,
|
|
33
|
+
data_path_glob=self.data_path_glob,
|
|
34
|
+
project_dir=self.project_dir,
|
|
35
|
+
all_files=all_files,
|
|
36
|
+
)
|
|
37
|
+
logger.info("Parsing Done!")
|
autorag/schema/base.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
from abc import ABCMeta, abstractmethod
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from typing import Union
|
|
4
|
+
|
|
5
|
+
import pandas as pd
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class BaseModule(metaclass=ABCMeta):
|
|
9
|
+
@abstractmethod
|
|
10
|
+
def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
|
|
11
|
+
pass
|
|
12
|
+
|
|
13
|
+
@abstractmethod
|
|
14
|
+
def _pure(self, *args, **kwargs):
|
|
15
|
+
pass
|
|
16
|
+
|
|
17
|
+
@classmethod
|
|
18
|
+
def run_evaluator(
|
|
19
|
+
cls,
|
|
20
|
+
project_dir: Union[str, Path],
|
|
21
|
+
previous_result: pd.DataFrame,
|
|
22
|
+
*args,
|
|
23
|
+
**kwargs,
|
|
24
|
+
):
|
|
25
|
+
instance = cls(project_dir, *args, **kwargs)
|
|
26
|
+
result = instance.pure(previous_result, *args, **kwargs)
|
|
27
|
+
del instance
|
|
28
|
+
return result
|
|
29
|
+
|
|
30
|
+
@abstractmethod
|
|
31
|
+
def cast_to_run(self, previous_result: pd.DataFrame, *args, **kwargs):
|
|
32
|
+
"""
|
|
33
|
+
This function is for cast function (a.k.a decorator) only for pure function in the whole node.
|
|
34
|
+
"""
|
|
35
|
+
pass
|
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from typing import Optional, List, Dict, Callable, Any, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import pandas as pd
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass
|
|
9
|
+
class MetricInput:
|
|
10
|
+
query: Optional[str] = None
|
|
11
|
+
queries: Optional[List[str]] = None
|
|
12
|
+
retrieval_gt_contents: Optional[List[List[str]]] = None
|
|
13
|
+
retrieved_contents: Optional[List[str]] = None
|
|
14
|
+
retrieval_gt: Optional[List[List[str]]] = None
|
|
15
|
+
retrieved_ids: Optional[List[str]] = None
|
|
16
|
+
prompt: Optional[str] = None
|
|
17
|
+
generated_texts: Optional[str] = None
|
|
18
|
+
generation_gt: Optional[List[str]] = None
|
|
19
|
+
generated_log_probs: Optional[List[float]] = None
|
|
20
|
+
|
|
21
|
+
def is_fields_notnone(self, fields_to_check: List[str]) -> bool:
|
|
22
|
+
for field in fields_to_check:
|
|
23
|
+
actual_value = getattr(self, field)
|
|
24
|
+
|
|
25
|
+
if actual_value is None:
|
|
26
|
+
return False
|
|
27
|
+
|
|
28
|
+
try:
|
|
29
|
+
if not type_checks.get(type(actual_value), lambda _: False)(
|
|
30
|
+
actual_value
|
|
31
|
+
):
|
|
32
|
+
return False
|
|
33
|
+
except Exception:
|
|
34
|
+
return False
|
|
35
|
+
|
|
36
|
+
return True
|
|
37
|
+
|
|
38
|
+
@classmethod
|
|
39
|
+
def from_dataframe(cls, qa_data: pd.DataFrame) -> List["MetricInput"]:
|
|
40
|
+
"""
|
|
41
|
+
Convert a pandas DataFrame into a list of MetricInput instances.
|
|
42
|
+
qa_data: pd.DataFrame: qa_data DataFrame containing metric data.
|
|
43
|
+
|
|
44
|
+
:returns: List[MetricInput]: List of MetricInput objects created from DataFrame rows.
|
|
45
|
+
"""
|
|
46
|
+
instances = []
|
|
47
|
+
|
|
48
|
+
for _, row in qa_data.iterrows():
|
|
49
|
+
instance = cls()
|
|
50
|
+
|
|
51
|
+
for attr_name in cls.__annotations__:
|
|
52
|
+
if attr_name in row:
|
|
53
|
+
value = row[attr_name]
|
|
54
|
+
|
|
55
|
+
if isinstance(value, str):
|
|
56
|
+
setattr(
|
|
57
|
+
instance,
|
|
58
|
+
attr_name,
|
|
59
|
+
value.strip() if value.strip() != "" else None,
|
|
60
|
+
)
|
|
61
|
+
elif isinstance(value, list):
|
|
62
|
+
setattr(instance, attr_name, value if len(value) > 0 else None)
|
|
63
|
+
else:
|
|
64
|
+
setattr(instance, attr_name, value)
|
|
65
|
+
|
|
66
|
+
instances.append(instance)
|
|
67
|
+
|
|
68
|
+
return instances
|
|
69
|
+
|
|
70
|
+
@staticmethod
|
|
71
|
+
def _check_list(lst_or_arr: Union[List[Any], np.ndarray]) -> bool:
|
|
72
|
+
if isinstance(lst_or_arr, np.ndarray):
|
|
73
|
+
lst_or_arr = lst_or_arr.flatten().tolist()
|
|
74
|
+
|
|
75
|
+
if len(lst_or_arr) == 0:
|
|
76
|
+
return False
|
|
77
|
+
|
|
78
|
+
for item in lst_or_arr:
|
|
79
|
+
if item is None:
|
|
80
|
+
return False
|
|
81
|
+
|
|
82
|
+
item_type = type(item)
|
|
83
|
+
|
|
84
|
+
if item_type in type_checks:
|
|
85
|
+
if not type_checks[item_type](item):
|
|
86
|
+
return False
|
|
87
|
+
else:
|
|
88
|
+
return False
|
|
89
|
+
|
|
90
|
+
return True
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
type_checks: Dict[type, Callable[[Any], bool]] = {
|
|
94
|
+
str: lambda x: len(x.strip()) > 0,
|
|
95
|
+
list: MetricInput._check_list,
|
|
96
|
+
np.ndarray: MetricInput._check_list,
|
|
97
|
+
int: lambda _: True,
|
|
98
|
+
float: lambda _: True,
|
|
99
|
+
}
|
autorag/schema/module.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
from copy import deepcopy
|
|
2
|
+
from dataclasses import dataclass, field
|
|
3
|
+
from typing import Callable, Dict
|
|
4
|
+
|
|
5
|
+
from autorag.support import get_support_modules
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass
|
|
9
|
+
class Module:
|
|
10
|
+
module_type: str
|
|
11
|
+
module_param: Dict
|
|
12
|
+
module: Callable = field(init=False)
|
|
13
|
+
|
|
14
|
+
def __post_init__(self):
|
|
15
|
+
self.module = get_support_modules(self.module_type)
|
|
16
|
+
if self.module is None:
|
|
17
|
+
raise ValueError(f"Module type {self.module_type} is not supported.")
|
|
18
|
+
|
|
19
|
+
@classmethod
|
|
20
|
+
def from_dict(cls, module_dict: Dict) -> "Module":
|
|
21
|
+
_module_dict = deepcopy(module_dict)
|
|
22
|
+
module_type = _module_dict.pop("module_type")
|
|
23
|
+
module_params = _module_dict
|
|
24
|
+
return cls(module_type, module_params)
|
autorag/schema/node.py
ADDED
|
@@ -0,0 +1,144 @@
|
|
|
1
|
+
import itertools
|
|
2
|
+
import logging
|
|
3
|
+
from copy import deepcopy
|
|
4
|
+
from dataclasses import dataclass, field
|
|
5
|
+
from typing import Dict, List, Callable, Tuple, Any
|
|
6
|
+
|
|
7
|
+
import pandas as pd
|
|
8
|
+
|
|
9
|
+
from autorag.schema.module import Module
|
|
10
|
+
from autorag.support import get_support_nodes
|
|
11
|
+
from autorag.utils.util import make_combinations, explode, find_key_values
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger("AutoRAG")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@dataclass
|
|
17
|
+
class Node:
|
|
18
|
+
node_type: str
|
|
19
|
+
strategy: Dict
|
|
20
|
+
node_params: Dict
|
|
21
|
+
modules: List[Module]
|
|
22
|
+
run_node: Callable = field(init=False)
|
|
23
|
+
|
|
24
|
+
def __post_init__(self):
|
|
25
|
+
self.run_node = get_support_nodes(self.node_type)
|
|
26
|
+
if self.run_node is None:
|
|
27
|
+
raise ValueError(f"Node type {self.node_type} is not supported.")
|
|
28
|
+
|
|
29
|
+
def get_param_combinations(self) -> Tuple[List[Callable], List[Dict]]:
|
|
30
|
+
"""
|
|
31
|
+
This method returns a combination of module and node parameters, also corresponding modules.
|
|
32
|
+
|
|
33
|
+
:return: Each module and its module parameters.
|
|
34
|
+
:rtype: Tuple[List[Callable], List[Dict]]
|
|
35
|
+
"""
|
|
36
|
+
|
|
37
|
+
def make_single_combination(module: Module) -> List[Dict]:
|
|
38
|
+
input_dict = {**self.node_params, **module.module_param}
|
|
39
|
+
return make_combinations(input_dict)
|
|
40
|
+
|
|
41
|
+
combinations = list(map(make_single_combination, self.modules))
|
|
42
|
+
module_list, combination_list = explode(self.modules, combinations)
|
|
43
|
+
return list(map(lambda x: x.module, module_list)), combination_list
|
|
44
|
+
|
|
45
|
+
@classmethod
|
|
46
|
+
def from_dict(cls, node_dict: Dict) -> "Node":
|
|
47
|
+
_node_dict = deepcopy(node_dict)
|
|
48
|
+
node_type = _node_dict.pop("node_type")
|
|
49
|
+
strategy = _node_dict.pop("strategy")
|
|
50
|
+
modules = list(map(lambda x: Module.from_dict(x), _node_dict.pop("modules")))
|
|
51
|
+
node_params = _node_dict
|
|
52
|
+
return cls(node_type, strategy, node_params, modules)
|
|
53
|
+
|
|
54
|
+
def run(self, previous_result: pd.DataFrame, node_line_dir: str) -> pd.DataFrame:
|
|
55
|
+
logger.info(f"Running node {self.node_type}...")
|
|
56
|
+
input_modules, input_params = self.get_param_combinations()
|
|
57
|
+
return self.run_node(
|
|
58
|
+
modules=input_modules,
|
|
59
|
+
module_params=input_params,
|
|
60
|
+
previous_result=previous_result,
|
|
61
|
+
node_line_dir=node_line_dir,
|
|
62
|
+
strategies=self.strategy,
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def extract_values(node: Node, key: str) -> List[str]:
|
|
67
|
+
"""
|
|
68
|
+
This function extract values from node's modules' module_param.
|
|
69
|
+
|
|
70
|
+
:param node: The node you want to extract values from.
|
|
71
|
+
:param key: The key of module_param that you want to extract.
|
|
72
|
+
:return: The list of extracted values.
|
|
73
|
+
It removes duplicated elements automatically.
|
|
74
|
+
"""
|
|
75
|
+
|
|
76
|
+
def extract_module_values(module: Module):
|
|
77
|
+
if key not in module.module_param:
|
|
78
|
+
return []
|
|
79
|
+
value = module.module_param[key]
|
|
80
|
+
if isinstance(value, str) or isinstance(value, int):
|
|
81
|
+
return [value]
|
|
82
|
+
elif isinstance(value, list):
|
|
83
|
+
return value
|
|
84
|
+
else:
|
|
85
|
+
raise ValueError(f"{key} must be str,list or int, but got {type(value)}")
|
|
86
|
+
|
|
87
|
+
values = list(map(extract_module_values, node.modules))
|
|
88
|
+
return list(set(list(itertools.chain.from_iterable(values))))
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def extract_values_from_nodes(nodes: List[Node], key: str) -> List[str]:
|
|
92
|
+
"""
|
|
93
|
+
This function extract values from nodes' modules' module_param.
|
|
94
|
+
|
|
95
|
+
:param nodes: The nodes you want to extract values from.
|
|
96
|
+
:param key: The key of module_param that you want to extract.
|
|
97
|
+
:return: The list of extracted values.
|
|
98
|
+
It removes duplicated elements automatically.
|
|
99
|
+
"""
|
|
100
|
+
values = list(map(lambda node: extract_values(node, key), nodes))
|
|
101
|
+
return list(set(list(itertools.chain.from_iterable(values))))
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def extract_values_from_nodes_strategy(nodes: List[Node], key: str) -> List[Any]:
|
|
105
|
+
"""
|
|
106
|
+
This function extract values from nodes' strategy.
|
|
107
|
+
|
|
108
|
+
:param nodes: The nodes you want to extract values from.
|
|
109
|
+
:param key: The key string that you want to extract.
|
|
110
|
+
:return: The list of extracted values.
|
|
111
|
+
It removes duplicated elements automatically.
|
|
112
|
+
"""
|
|
113
|
+
values = []
|
|
114
|
+
for node in nodes:
|
|
115
|
+
value_list = find_key_values(node.strategy, key)
|
|
116
|
+
if value_list:
|
|
117
|
+
values.extend(value_list)
|
|
118
|
+
return values
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def module_type_exists(nodes: List[Node], module_type: str) -> bool:
|
|
122
|
+
"""
|
|
123
|
+
This function check if the module type exists in the nodes.
|
|
124
|
+
|
|
125
|
+
:param nodes: The nodes you want to check.
|
|
126
|
+
:param module_type: The module type you want to check.
|
|
127
|
+
:return: True if the module type exists in the nodes.
|
|
128
|
+
"""
|
|
129
|
+
return any(
|
|
130
|
+
list(
|
|
131
|
+
map(
|
|
132
|
+
lambda node: any(
|
|
133
|
+
list(
|
|
134
|
+
map(
|
|
135
|
+
lambda module: module.module_type.lower()
|
|
136
|
+
== module_type.lower(),
|
|
137
|
+
node.modules,
|
|
138
|
+
)
|
|
139
|
+
)
|
|
140
|
+
),
|
|
141
|
+
nodes,
|
|
142
|
+
)
|
|
143
|
+
)
|
|
144
|
+
)
|