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.
Files changed (184) hide show
  1. autorag/__init__.py +82 -0
  2. autorag/chunker.py +51 -0
  3. autorag/cli.py +209 -0
  4. autorag/dashboard.py +199 -0
  5. autorag/data/__init__.py +109 -0
  6. autorag/data/chunk/__init__.py +2 -0
  7. autorag/data/chunk/base.py +128 -0
  8. autorag/data/chunk/langchain_chunk.py +76 -0
  9. autorag/data/chunk/llama_index_chunk.py +96 -0
  10. autorag/data/chunk/run.py +38 -0
  11. autorag/data/legacy/__init__.py +0 -0
  12. autorag/data/legacy/corpus/__init__.py +2 -0
  13. autorag/data/legacy/corpus/langchain.py +47 -0
  14. autorag/data/legacy/corpus/llama_index.py +93 -0
  15. autorag/data/legacy/qacreation/__init__.py +6 -0
  16. autorag/data/legacy/qacreation/base.py +239 -0
  17. autorag/data/legacy/qacreation/llama_index.py +253 -0
  18. autorag/data/legacy/qacreation/llama_index_default_prompt.txt +54 -0
  19. autorag/data/legacy/qacreation/ragas.py +75 -0
  20. autorag/data/legacy/qacreation/simple.py +99 -0
  21. autorag/data/parse/__init__.py +1 -0
  22. autorag/data/parse/base.py +79 -0
  23. autorag/data/parse/clova.py +194 -0
  24. autorag/data/parse/langchain_parse.py +87 -0
  25. autorag/data/parse/llamaparse.py +126 -0
  26. autorag/data/parse/run.py +141 -0
  27. autorag/data/parse/table_hybrid_parse.py +134 -0
  28. autorag/data/qa/__init__.py +3 -0
  29. autorag/data/qa/evolve/__init__.py +0 -0
  30. autorag/data/qa/evolve/llama_index_query_evolve.py +64 -0
  31. autorag/data/qa/evolve/openai_query_evolve.py +81 -0
  32. autorag/data/qa/evolve/prompt.py +288 -0
  33. autorag/data/qa/extract_evidence.py +1 -0
  34. autorag/data/qa/filter/__init__.py +0 -0
  35. autorag/data/qa/filter/dontknow.py +117 -0
  36. autorag/data/qa/filter/passage_dependency.py +88 -0
  37. autorag/data/qa/filter/prompt.py +73 -0
  38. autorag/data/qa/generation_gt/__init__.py +0 -0
  39. autorag/data/qa/generation_gt/base.py +16 -0
  40. autorag/data/qa/generation_gt/llama_index_gen_gt.py +41 -0
  41. autorag/data/qa/generation_gt/openai_gen_gt.py +84 -0
  42. autorag/data/qa/generation_gt/prompt.py +27 -0
  43. autorag/data/qa/query/__init__.py +0 -0
  44. autorag/data/qa/query/llama_gen_query.py +82 -0
  45. autorag/data/qa/query/openai_gen_query.py +95 -0
  46. autorag/data/qa/query/prompt.py +201 -0
  47. autorag/data/qa/sample.py +26 -0
  48. autorag/data/qa/schema.py +322 -0
  49. autorag/data/utils/__init__.py +0 -0
  50. autorag/data/utils/util.py +103 -0
  51. autorag/deploy/__init__.py +9 -0
  52. autorag/deploy/api.py +303 -0
  53. autorag/deploy/base.py +235 -0
  54. autorag/deploy/gradio.py +74 -0
  55. autorag/deploy/swagger.yml +202 -0
  56. autorag/embedding/__init__.py +0 -0
  57. autorag/embedding/base.py +144 -0
  58. autorag/embedding/vllm.py +256 -0
  59. autorag/evaluation/__init__.py +3 -0
  60. autorag/evaluation/generation.py +88 -0
  61. autorag/evaluation/metric/__init__.py +22 -0
  62. autorag/evaluation/metric/deepeval_prompt.py +322 -0
  63. autorag/evaluation/metric/g_eval_prompts/coh_detailed.txt +32 -0
  64. autorag/evaluation/metric/g_eval_prompts/con_detailed.txt +33 -0
  65. autorag/evaluation/metric/g_eval_prompts/flu_detailed.txt +26 -0
  66. autorag/evaluation/metric/g_eval_prompts/rel_detailed.txt +33 -0
  67. autorag/evaluation/metric/generation.py +504 -0
  68. autorag/evaluation/metric/retrieval.py +115 -0
  69. autorag/evaluation/metric/retrieval_contents.py +65 -0
  70. autorag/evaluation/metric/util.py +88 -0
  71. autorag/evaluation/retrieval.py +83 -0
  72. autorag/evaluation/retrieval_contents.py +65 -0
  73. autorag/evaluation/util.py +43 -0
  74. autorag/evaluator.py +559 -0
  75. autorag/node_line.py +65 -0
  76. autorag/nodes/__init__.py +0 -0
  77. autorag/nodes/generator/__init__.py +4 -0
  78. autorag/nodes/generator/base.py +103 -0
  79. autorag/nodes/generator/llama_index_llm.py +169 -0
  80. autorag/nodes/generator/openai_llm.py +329 -0
  81. autorag/nodes/generator/run.py +148 -0
  82. autorag/nodes/generator/vllm.py +147 -0
  83. autorag/nodes/generator/vllm_api.py +191 -0
  84. autorag/nodes/hybridretrieval/__init__.py +2 -0
  85. autorag/nodes/hybridretrieval/base.py +58 -0
  86. autorag/nodes/hybridretrieval/hybrid_cc.py +227 -0
  87. autorag/nodes/hybridretrieval/hybrid_rrf.py +149 -0
  88. autorag/nodes/hybridretrieval/run.py +137 -0
  89. autorag/nodes/lexicalretrieval/__init__.py +1 -0
  90. autorag/nodes/lexicalretrieval/bm25.py +381 -0
  91. autorag/nodes/lexicalretrieval/run.py +148 -0
  92. autorag/nodes/passageaugmenter/__init__.py +2 -0
  93. autorag/nodes/passageaugmenter/base.py +76 -0
  94. autorag/nodes/passageaugmenter/pass_passage_augmenter.py +43 -0
  95. autorag/nodes/passageaugmenter/prev_next_augmenter.py +155 -0
  96. autorag/nodes/passageaugmenter/run.py +131 -0
  97. autorag/nodes/passagecompressor/__init__.py +4 -0
  98. autorag/nodes/passagecompressor/base.py +78 -0
  99. autorag/nodes/passagecompressor/longllmlingua.py +115 -0
  100. autorag/nodes/passagecompressor/pass_compressor.py +16 -0
  101. autorag/nodes/passagecompressor/refine.py +54 -0
  102. autorag/nodes/passagecompressor/run.py +186 -0
  103. autorag/nodes/passagecompressor/tree_summarize.py +56 -0
  104. autorag/nodes/passagefilter/__init__.py +6 -0
  105. autorag/nodes/passagefilter/base.py +40 -0
  106. autorag/nodes/passagefilter/pass_passage_filter.py +14 -0
  107. autorag/nodes/passagefilter/percentile_cutoff.py +58 -0
  108. autorag/nodes/passagefilter/recency.py +105 -0
  109. autorag/nodes/passagefilter/run.py +138 -0
  110. autorag/nodes/passagefilter/similarity_percentile_cutoff.py +134 -0
  111. autorag/nodes/passagefilter/similarity_threshold_cutoff.py +112 -0
  112. autorag/nodes/passagefilter/threshold_cutoff.py +78 -0
  113. autorag/nodes/passagereranker/__init__.py +16 -0
  114. autorag/nodes/passagereranker/base.py +44 -0
  115. autorag/nodes/passagereranker/cohere.py +118 -0
  116. autorag/nodes/passagereranker/colbert.py +213 -0
  117. autorag/nodes/passagereranker/flag_embedding.py +112 -0
  118. autorag/nodes/passagereranker/flag_embedding_llm.py +101 -0
  119. autorag/nodes/passagereranker/flashrank.py +245 -0
  120. autorag/nodes/passagereranker/jina.py +115 -0
  121. autorag/nodes/passagereranker/koreranker.py +136 -0
  122. autorag/nodes/passagereranker/mixedbreadai.py +126 -0
  123. autorag/nodes/passagereranker/monot5.py +190 -0
  124. autorag/nodes/passagereranker/openvino.py +191 -0
  125. autorag/nodes/passagereranker/pass_reranker.py +31 -0
  126. autorag/nodes/passagereranker/rankgpt.py +170 -0
  127. autorag/nodes/passagereranker/run.py +145 -0
  128. autorag/nodes/passagereranker/sentence_transformer.py +129 -0
  129. autorag/nodes/passagereranker/tart/__init__.py +1 -0
  130. autorag/nodes/passagereranker/tart/modeling_enc_t5.py +152 -0
  131. autorag/nodes/passagereranker/tart/tart.py +139 -0
  132. autorag/nodes/passagereranker/tart/tokenization_enc_t5.py +112 -0
  133. autorag/nodes/passagereranker/time_reranker.py +72 -0
  134. autorag/nodes/passagereranker/upr.py +160 -0
  135. autorag/nodes/passagereranker/voyageai.py +109 -0
  136. autorag/nodes/promptmaker/__init__.py +12 -0
  137. autorag/nodes/promptmaker/base.py +32 -0
  138. autorag/nodes/promptmaker/chat_fstring.py +73 -0
  139. autorag/nodes/promptmaker/fstring.py +49 -0
  140. autorag/nodes/promptmaker/long_context_reorder.py +83 -0
  141. autorag/nodes/promptmaker/run.py +283 -0
  142. autorag/nodes/promptmaker/window_replacement.py +85 -0
  143. autorag/nodes/queryexpansion/__init__.py +4 -0
  144. autorag/nodes/queryexpansion/base.py +62 -0
  145. autorag/nodes/queryexpansion/hyde.py +43 -0
  146. autorag/nodes/queryexpansion/multi_query_expansion.py +57 -0
  147. autorag/nodes/queryexpansion/pass_query_expansion.py +22 -0
  148. autorag/nodes/queryexpansion/query_decompose.py +111 -0
  149. autorag/nodes/queryexpansion/run.py +308 -0
  150. autorag/nodes/retrieval/__init__.py +0 -0
  151. autorag/nodes/retrieval/base.py +127 -0
  152. autorag/nodes/retrieval/run_util.py +152 -0
  153. autorag/nodes/semanticretrieval/__init__.py +1 -0
  154. autorag/nodes/semanticretrieval/run.py +148 -0
  155. autorag/nodes/semanticretrieval/vectordb.py +339 -0
  156. autorag/nodes/util.py +16 -0
  157. autorag/parser.py +37 -0
  158. autorag/schema/__init__.py +3 -0
  159. autorag/schema/base.py +35 -0
  160. autorag/schema/metricinput.py +99 -0
  161. autorag/schema/module.py +24 -0
  162. autorag/schema/node.py +144 -0
  163. autorag/strategy.py +165 -0
  164. autorag/support.py +235 -0
  165. autorag/utils/__init__.py +8 -0
  166. autorag/utils/cast.py +45 -0
  167. autorag/utils/preprocess.py +149 -0
  168. autorag/utils/util.py +759 -0
  169. autorag/validator.py +98 -0
  170. autorag/vectordb/__init__.py +75 -0
  171. autorag/vectordb/base.py +73 -0
  172. autorag/vectordb/chroma.py +118 -0
  173. autorag/vectordb/couchbase.py +239 -0
  174. autorag/vectordb/milvus.py +169 -0
  175. autorag/vectordb/pinecone.py +121 -0
  176. autorag/vectordb/qdrant.py +155 -0
  177. autorag/vectordb/weaviate.py +184 -0
  178. autorag/web.py +81 -0
  179. autorag-0.0.0.dist-info/METADATA +780 -0
  180. autorag-0.0.0.dist-info/RECORD +184 -0
  181. autorag-0.0.0.dist-info/WHEEL +5 -0
  182. autorag-0.0.0.dist-info/entry_points.txt +2 -0
  183. autorag-0.0.0.dist-info/licenses/LICENSE +201 -0
  184. 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!")
@@ -0,0 +1,3 @@
1
+ from .module import Module
2
+ from .node import Node
3
+ from .base import BaseModule
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
+ }
@@ -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
+ )