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,213 @@
1
+ from typing import List, Tuple
2
+
3
+ import numpy as np
4
+ import pandas as pd
5
+
6
+ from autorag.nodes.passagereranker.base import BasePassageReranker
7
+ from autorag.utils.util import (
8
+ flatten_apply,
9
+ sort_by_scores,
10
+ select_top_k,
11
+ pop_params,
12
+ result_to_dataframe,
13
+ empty_cuda_cache,
14
+ )
15
+
16
+
17
+ class ColbertReranker(BasePassageReranker):
18
+ def __init__(
19
+ self,
20
+ project_dir: str,
21
+ model_name: str = "colbert-ir/colbertv2.0",
22
+ *args,
23
+ **kwargs,
24
+ ):
25
+ """
26
+ Initialize a colbert rerank model for reranking.
27
+
28
+ :param project_dir: The project directory
29
+ :param model_name: The model name for Colbert rerank.
30
+ You can choose a colbert model for reranking.
31
+ The default is "colbert-ir/colbertv2.0".
32
+ :param kwargs: Extra parameter for the model.
33
+ """
34
+ super().__init__(project_dir)
35
+ try:
36
+ import torch
37
+ from transformers import AutoModel, AutoTokenizer
38
+ except ImportError:
39
+ raise ImportError(
40
+ "Pytorch is not installed. Please install pytorch to use Colbert reranker."
41
+ )
42
+ self.device = "cuda" if torch.cuda.is_available() else "cpu"
43
+ model_params = pop_params(AutoModel.from_pretrained, kwargs)
44
+ self.model = AutoModel.from_pretrained(model_name, **model_params).to(
45
+ self.device
46
+ )
47
+ self.tokenizer = AutoTokenizer.from_pretrained(model_name)
48
+
49
+ def __del__(self):
50
+ del self.model
51
+ empty_cuda_cache()
52
+ super().__del__()
53
+
54
+ @result_to_dataframe(["retrieved_contents", "retrieved_ids", "retrieve_scores"])
55
+ def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
56
+ queries, contents, _, ids = self.cast_to_run(previous_result)
57
+ top_k = kwargs.pop("top_k")
58
+ batch = kwargs.pop("batch", 64)
59
+ return self._pure(queries, contents, ids, top_k, batch)
60
+
61
+ def _pure(
62
+ self,
63
+ queries: List[str],
64
+ contents_list: List[List[str]],
65
+ ids_list: List[List[str]],
66
+ top_k: int,
67
+ batch: int = 64,
68
+ ) -> Tuple[List[List[str]], List[List[str]], List[List[float]]]:
69
+ """
70
+ Rerank a list of contents with Colbert rerank models.
71
+ You can get more information about a Colbert model at https://huggingface.co/colbert-ir/colbertv2.0.
72
+ It uses BERT-based model, so recommend using CUDA gpu for faster reranking.
73
+
74
+ :param queries: The list of queries to use for reranking
75
+ :param contents_list: The list of lists of contents to rerank
76
+ :param ids_list: The list of lists of ids retrieved from the initial ranking
77
+ :param top_k: The number of passages to be retrieved
78
+ :param batch: The number of queries to be processed in a batch
79
+ Default is 64.
80
+
81
+ :return: Tuple of lists containing the reranked contents, ids, and scores
82
+ """
83
+
84
+ # get query and content embeddings
85
+ query_embedding_list = get_colbert_embedding_batch(
86
+ queries, self.model, self.tokenizer, batch
87
+ )
88
+ content_embedding_list = flatten_apply(
89
+ get_colbert_embedding_batch,
90
+ contents_list,
91
+ model=self.model,
92
+ tokenizer=self.tokenizer,
93
+ batch_size=batch,
94
+ )
95
+ df = pd.DataFrame(
96
+ {
97
+ "ids": ids_list,
98
+ "query_embedding": query_embedding_list,
99
+ "contents": contents_list,
100
+ "content_embedding": content_embedding_list,
101
+ }
102
+ )
103
+ temp_df = df.explode("content_embedding")
104
+ temp_df["score"] = temp_df.apply(
105
+ lambda x: get_colbert_score(x["query_embedding"], x["content_embedding"]),
106
+ axis=1,
107
+ )
108
+ df["scores"] = (
109
+ temp_df.groupby(level=0, sort=False)["score"].apply(list).tolist()
110
+ )
111
+ df[["contents", "ids", "scores"]] = df.apply(
112
+ sort_by_scores, axis=1, result_type="expand"
113
+ )
114
+ results = select_top_k(df, ["contents", "ids", "scores"], top_k)
115
+
116
+ return (
117
+ results["contents"].tolist(),
118
+ results["ids"].tolist(),
119
+ results["scores"].tolist(),
120
+ )
121
+
122
+
123
+ def get_colbert_embedding_batch(
124
+ input_strings: List[str], model, tokenizer, batch_size: int
125
+ ) -> List[np.array]:
126
+ try:
127
+ import torch
128
+ except ImportError:
129
+ raise ImportError(
130
+ "Pytorch is not installed. Please install pytorch to use Colbert reranker."
131
+ )
132
+ encoding = tokenizer(
133
+ input_strings,
134
+ return_tensors="pt",
135
+ padding=True,
136
+ truncation=True,
137
+ max_length=model.config.max_position_embeddings,
138
+ )
139
+
140
+ input_batches = slice_tokenizer_result(encoding, batch_size)
141
+ result_embedding = []
142
+ with torch.no_grad():
143
+ for encoding_batch in input_batches:
144
+ result_embedding.append(model(**encoding_batch).last_hidden_state)
145
+ total_tensor = torch.cat(
146
+ result_embedding, dim=0
147
+ ) # shape [batch_size, token_length, embedding_dim]
148
+ tensor_results = list(total_tensor.chunk(total_tensor.size()[0]))
149
+
150
+ if torch.cuda.is_available():
151
+ return list(map(lambda x: x.detach().cpu().numpy(), tensor_results))
152
+ else:
153
+ return list(map(lambda x: x.detach().numpy(), tensor_results))
154
+
155
+
156
+ def slice_tokenizer_result(tokenizer_output, batch_size):
157
+ input_ids_batches = slice_tensor(tokenizer_output["input_ids"], batch_size)
158
+ attention_mask_batches = slice_tensor(
159
+ tokenizer_output["attention_mask"], batch_size
160
+ )
161
+ token_type_ids_batches = slice_tensor(
162
+ tokenizer_output.get("token_type_ids", None), batch_size
163
+ )
164
+ return [
165
+ {
166
+ "input_ids": input_ids,
167
+ "attention_mask": attention_mask,
168
+ "token_type_ids": token_type_ids,
169
+ }
170
+ for input_ids, attention_mask, token_type_ids in zip(
171
+ input_ids_batches, attention_mask_batches, token_type_ids_batches
172
+ )
173
+ ]
174
+
175
+
176
+ def slice_tensor(input_tensor, batch_size):
177
+ try:
178
+ import torch
179
+ except ImportError:
180
+ raise ImportError(
181
+ "Pytorch is not installed. Please install pytorch to use Colbert reranker."
182
+ )
183
+ # Calculate the number of full batches
184
+ num_full_batches = input_tensor.size(0) // batch_size
185
+
186
+ # Slice the tensor into batches
187
+ tensor_list = [
188
+ input_tensor[i * batch_size : (i + 1) * batch_size]
189
+ for i in range(num_full_batches)
190
+ ]
191
+
192
+ # Handle the last batch if it's smaller than batch_size
193
+ remainder = input_tensor.size(0) % batch_size
194
+ if remainder:
195
+ tensor_list.append(input_tensor[-remainder:])
196
+
197
+ device = "cuda" if torch.cuda.is_available() else "cpu"
198
+ tensor_list = list(map(lambda x: x.to(device), tensor_list))
199
+
200
+ return tensor_list
201
+
202
+
203
+ def get_colbert_score(query_embedding: np.array, content_embedding: np.array) -> float:
204
+ if query_embedding.ndim == 3 and content_embedding.ndim == 3:
205
+ query_embedding = query_embedding.reshape(-1, query_embedding.shape[-1])
206
+ content_embedding = content_embedding.reshape(-1, content_embedding.shape[-1])
207
+
208
+ sim_matrix = np.dot(query_embedding, content_embedding.T) / (
209
+ np.linalg.norm(query_embedding, axis=1)[:, np.newaxis]
210
+ * np.linalg.norm(content_embedding, axis=1)
211
+ )
212
+ max_sim_scores = np.max(sim_matrix, axis=1)
213
+ return float(np.mean(max_sim_scores))
@@ -0,0 +1,112 @@
1
+ from typing import List, Tuple, Iterable
2
+
3
+ import pandas as pd
4
+
5
+ from autorag.nodes.passagereranker.base import BasePassageReranker
6
+ from autorag.utils.util import (
7
+ make_batch,
8
+ sort_by_scores,
9
+ flatten_apply,
10
+ select_top_k,
11
+ pop_params,
12
+ result_to_dataframe,
13
+ empty_cuda_cache,
14
+ )
15
+
16
+
17
+ class FlagEmbeddingReranker(BasePassageReranker):
18
+ def __init__(
19
+ self, project_dir, model_name: str = "BAAI/bge-reranker-large", *args, **kwargs
20
+ ):
21
+ """
22
+ Initialize the FlagEmbeddingReranker module.
23
+
24
+ :param project_dir: The project directory.
25
+ :param model_name: The name of the BAAI Reranker normal-model name.
26
+ Default is "BAAI/bge-reranker-large"
27
+ :param kwargs: Extra parameter for FlagEmbedding.FlagReranker
28
+ """
29
+ super().__init__(project_dir)
30
+ try:
31
+ from FlagEmbedding import FlagReranker
32
+ except ImportError:
33
+ raise ImportError(
34
+ "FlagEmbeddingReranker requires the 'FlagEmbedding' package to be installed."
35
+ )
36
+ model_params = pop_params(FlagReranker.__init__, kwargs)
37
+ model_params.pop("model_name_or_path", None)
38
+ self.model = FlagReranker(model_name_or_path=model_name, **model_params)
39
+
40
+ def __del__(self):
41
+ del self.model
42
+ empty_cuda_cache()
43
+ super().__del__()
44
+
45
+ @result_to_dataframe(["retrieved_contents", "retrieved_ids", "retrieve_scores"])
46
+ def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
47
+ queries, contents, _, ids = self.cast_to_run(previous_result)
48
+ top_k = kwargs.pop("top_k")
49
+ batch = kwargs.pop("batch", 64)
50
+ return self._pure(queries, contents, ids, top_k, batch)
51
+
52
+ def _pure(
53
+ self,
54
+ queries: List[str],
55
+ contents_list: List[List[str]],
56
+ ids_list: List[List[str]],
57
+ top_k: int,
58
+ batch: int = 64,
59
+ ) -> Tuple[List[List[str]], List[List[str]], List[List[float]]]:
60
+ """
61
+ Rerank a list of contents based on their relevance to a query using BAAI normal-Reranker model.
62
+
63
+ :param queries: The list of queries to use for reranking
64
+ :param contents_list: The list of lists of contents to rerank
65
+ :param ids_list: The list of lists of ids retrieved from the initial ranking
66
+ :param top_k: The number of passages to be retrieved
67
+ :param batch: The number of queries to be processed in a batch
68
+ Default is 64.
69
+ :return: Tuple of lists containing the reranked contents, ids, and scores
70
+ """
71
+ nested_list = [
72
+ list(map(lambda x: [query, x], content_list))
73
+ for query, content_list in zip(queries, contents_list)
74
+ ]
75
+ rerank_scores = flatten_apply(
76
+ flag_embedding_run_model, nested_list, model=self.model, batch_size=batch
77
+ )
78
+
79
+ df = pd.DataFrame(
80
+ {
81
+ "contents": contents_list,
82
+ "ids": ids_list,
83
+ "scores": rerank_scores,
84
+ }
85
+ )
86
+ df[["contents", "ids", "scores"]] = df.apply(
87
+ sort_by_scores, axis=1, result_type="expand"
88
+ )
89
+ results = select_top_k(df, ["contents", "ids", "scores"], top_k)
90
+
91
+ return (
92
+ results["contents"].tolist(),
93
+ results["ids"].tolist(),
94
+ results["scores"].tolist(),
95
+ )
96
+
97
+
98
+ def flag_embedding_run_model(input_texts, model, batch_size: int):
99
+ try:
100
+ import torch
101
+ except ImportError:
102
+ raise ImportError("FlagEmbeddingReranker requires PyTorch to be installed.")
103
+ batch_input_texts = make_batch(input_texts, batch_size)
104
+ results = []
105
+ for batch_texts in batch_input_texts:
106
+ with torch.no_grad():
107
+ pred_scores = model.compute_score(sentence_pairs=batch_texts)
108
+ if not isinstance(pred_scores, Iterable):
109
+ results.append(pred_scores)
110
+ else:
111
+ results.extend(pred_scores)
112
+ return results
@@ -0,0 +1,101 @@
1
+ from typing import List, Tuple
2
+
3
+ import pandas as pd
4
+
5
+ from autorag.nodes.passagereranker.base import BasePassageReranker
6
+ from autorag.nodes.passagereranker.flag_embedding import flag_embedding_run_model
7
+ from autorag.utils.util import (
8
+ flatten_apply,
9
+ sort_by_scores,
10
+ select_top_k,
11
+ pop_params,
12
+ result_to_dataframe,
13
+ empty_cuda_cache,
14
+ )
15
+
16
+
17
+ class FlagEmbeddingLLMReranker(BasePassageReranker):
18
+ def __init__(
19
+ self,
20
+ project_dir,
21
+ model_name: str = "BAAI/bge-reranker-v2-gemma",
22
+ *args,
23
+ **kwargs,
24
+ ):
25
+ """
26
+ Initialize the FlagEmbeddingReranker module.
27
+
28
+ :param project_dir: The project directory.
29
+ :param model_name: The name of the BAAI Reranker LLM-based-model name.
30
+ Default is "BAAI/bge-reranker-v2-gemma"
31
+ :param kwargs: Extra parameter for FlagEmbedding.FlagReranker
32
+ """
33
+ super().__init__(project_dir)
34
+ try:
35
+ from FlagEmbedding import FlagLLMReranker
36
+ except ImportError:
37
+ raise ImportError(
38
+ "FlagEmbeddingLLMReranker requires the 'FlagEmbedding' package to be installed."
39
+ )
40
+ model_params = pop_params(FlagLLMReranker.__init__, kwargs)
41
+ model_params.pop("model_name_or_path", None)
42
+ self.model = FlagLLMReranker(model_name_or_path=model_name, **model_params)
43
+
44
+ def __del__(self):
45
+ del self.model
46
+ empty_cuda_cache()
47
+ super().__del__()
48
+
49
+ @result_to_dataframe(["retrieved_contents", "retrieved_ids", "retrieve_scores"])
50
+ def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
51
+ queries, contents, _, ids = self.cast_to_run(previous_result)
52
+ top_k = kwargs.pop("top_k")
53
+ batch = kwargs.pop("batch", 64)
54
+ return self._pure(queries, contents, ids, top_k, batch)
55
+
56
+ def _pure(
57
+ self,
58
+ queries: List[str],
59
+ contents_list: List[List[str]],
60
+ ids_list: List[List[str]],
61
+ top_k: int,
62
+ batch: int = 64,
63
+ ) -> Tuple[List[List[str]], List[List[str]], List[List[float]]]:
64
+ """
65
+ Rerank a list of contents based on their relevance to a query using BAAI LLM-based-Reranker model.
66
+
67
+ :param queries: The list of queries to use for reranking
68
+ :param contents_list: The list of lists of contents to rerank
69
+ :param ids_list: The list of lists of ids retrieved from the initial ranking
70
+ :param top_k: The number of passages to be retrieved
71
+ :param batch: The number of queries to be processed in a batch
72
+ Default is 64.
73
+
74
+ :return: tuple of lists containing the reranked contents, ids, and scores
75
+ """
76
+
77
+ nested_list = [
78
+ list(map(lambda x: [query, x], content_list))
79
+ for query, content_list in zip(queries, contents_list)
80
+ ]
81
+ rerank_scores = flatten_apply(
82
+ flag_embedding_run_model, nested_list, model=self.model, batch_size=batch
83
+ )
84
+
85
+ df = pd.DataFrame(
86
+ {
87
+ "contents": contents_list,
88
+ "ids": ids_list,
89
+ "scores": rerank_scores,
90
+ }
91
+ )
92
+ df[["contents", "ids", "scores"]] = df.apply(
93
+ sort_by_scores, axis=1, result_type="expand"
94
+ )
95
+ results = select_top_k(df, ["contents", "ids", "scores"], top_k)
96
+
97
+ return (
98
+ results["contents"].tolist(),
99
+ results["ids"].tolist(),
100
+ results["scores"].tolist(),
101
+ )
@@ -0,0 +1,245 @@
1
+ import json
2
+ from pathlib import Path
3
+
4
+ import pandas as pd
5
+ import numpy as np
6
+ import os
7
+ import zipfile
8
+ import requests
9
+ from tqdm import tqdm
10
+ import collections
11
+ from typing import List, Dict, Tuple
12
+
13
+ from autorag.nodes.passagereranker.base import BasePassageReranker
14
+ from autorag.utils import result_to_dataframe
15
+ from autorag.utils.util import (
16
+ flatten_apply,
17
+ sort_by_scores,
18
+ select_top_k,
19
+ make_batch,
20
+ empty_cuda_cache,
21
+ )
22
+
23
+ model_url = "https://huggingface.co/prithivida/flashrank/resolve/main/{}.zip"
24
+
25
+ model_file_map = {
26
+ "ms-marco-TinyBERT-L-2-v2": "flashrank-TinyBERT-L-2-v2.onnx",
27
+ "ms-marco-MiniLM-L-12-v2": "flashrank-MiniLM-L-12-v2_Q.onnx",
28
+ "ms-marco-MultiBERT-L-12": "flashrank-MultiBERT-L12_Q.onnx",
29
+ "rank-T5-flan": "flashrank-rankt5_Q.onnx",
30
+ "ce-esci-MiniLM-L12-v2": "flashrank-ce-esci-MiniLM-L12-v2_Q.onnx",
31
+ "miniReranker_arabic_v1": "miniReranker_arabic_v1.onnx",
32
+ }
33
+
34
+
35
+ class FlashRankReranker(BasePassageReranker):
36
+ def __init__(
37
+ self, project_dir: str, model: str = "ms-marco-TinyBERT-L-2-v2", *args, **kwargs
38
+ ):
39
+ """
40
+ Initialize FlashRank rerank node.
41
+
42
+ :param project_dir: The project directory path.
43
+ :param model: The model name for FlashRank rerank.
44
+ You can get the list of available models from https://github.com/PrithivirajDamodaran/FlashRank.
45
+ Default is "ms-marco-TinyBERT-L-2-v2".
46
+ Not support “rank_zephyr_7b_v1_full” due to parallel inference issue.
47
+ :param kwargs: Extra arguments that are not affected
48
+ """
49
+ super().__init__(project_dir)
50
+ try:
51
+ from tokenizers import Tokenizer
52
+ except ImportError:
53
+ raise ImportError(
54
+ "Tokenizer is not installed. Please install tokenizers to use FlashRank reranker."
55
+ )
56
+
57
+ cache_dir = kwargs.pop("cache_dir", "/tmp")
58
+ max_length = kwargs.pop("max_length", 512)
59
+
60
+ self.cache_dir: Path = Path(cache_dir)
61
+ self.model_dir: Path = self.cache_dir / model
62
+ self._prepare_model_dir(model)
63
+ model_file = model_file_map[model]
64
+
65
+ try:
66
+ import onnxruntime as ort
67
+ except ImportError:
68
+ raise ImportError(
69
+ "onnxruntime is not installed. Please install onnxruntime to use FlashRank reranker."
70
+ )
71
+
72
+ self.session = ort.InferenceSession(str(self.model_dir / model_file))
73
+ self.tokenizer: Tokenizer = self._get_tokenizer(max_length)
74
+
75
+ def __del__(self):
76
+ del self.session
77
+ del self.tokenizer
78
+ empty_cuda_cache()
79
+ super().__del__()
80
+
81
+ def _prepare_model_dir(self, model_name: str):
82
+ if not self.cache_dir.exists():
83
+ self.cache_dir.mkdir(parents=True, exist_ok=True)
84
+
85
+ if not self.model_dir.exists():
86
+ self._download_model_files(model_name)
87
+
88
+ def _download_model_files(self, model_name: str):
89
+ local_zip_file = self.cache_dir / f"{model_name}.zip"
90
+ formatted_model_url = model_url.format(model_name)
91
+
92
+ with requests.get(formatted_model_url, stream=True) as r:
93
+ r.raise_for_status()
94
+ total_size = int(r.headers.get("content-length", 0))
95
+ with (
96
+ open(local_zip_file, "wb") as f,
97
+ tqdm(
98
+ desc=local_zip_file.name,
99
+ total=total_size,
100
+ unit="iB",
101
+ unit_scale=True,
102
+ unit_divisor=1024,
103
+ ) as bar,
104
+ ):
105
+ for chunk in r.iter_content(chunk_size=8192):
106
+ size = f.write(chunk)
107
+ bar.update(size)
108
+
109
+ with zipfile.ZipFile(local_zip_file, "r") as zip_ref:
110
+ zip_ref.extractall(self.cache_dir)
111
+ os.remove(local_zip_file)
112
+
113
+ def _get_tokenizer(self, max_length: int = 512):
114
+ try:
115
+ from tokenizers import AddedToken, Tokenizer
116
+ except ImportError:
117
+ raise ImportError(
118
+ "Pytorch is not installed. Please install pytorch to use FlashRank reranker."
119
+ )
120
+ config = json.load(open(str(self.model_dir / "config.json")))
121
+ tokenizer_config = json.load(
122
+ open(str(self.model_dir / "tokenizer_config.json"))
123
+ )
124
+ tokens_map = json.load(open(str(self.model_dir / "special_tokens_map.json")))
125
+ tokenizer = Tokenizer.from_file(str(self.model_dir / "tokenizer.json"))
126
+
127
+ tokenizer.enable_truncation(
128
+ max_length=min(tokenizer_config["model_max_length"], max_length)
129
+ )
130
+ tokenizer.enable_padding(
131
+ pad_id=config["pad_token_id"], pad_token=tokenizer_config["pad_token"]
132
+ )
133
+
134
+ for token in tokens_map.values():
135
+ if isinstance(token, str):
136
+ tokenizer.add_special_tokens([token])
137
+ elif isinstance(token, dict):
138
+ tokenizer.add_special_tokens([AddedToken(**token)])
139
+
140
+ vocab_file = self.model_dir / "vocab.txt"
141
+ if vocab_file.exists():
142
+ tokenizer.vocab = self._load_vocab(vocab_file)
143
+ tokenizer.ids_to_tokens = collections.OrderedDict(
144
+ [(ids, tok) for tok, ids in tokenizer.vocab.items()]
145
+ )
146
+ return tokenizer
147
+
148
+ def _load_vocab(self, vocab_file: Path) -> Dict[str, int]:
149
+ vocab = collections.OrderedDict()
150
+ with open(vocab_file, "r", encoding="utf-8") as reader:
151
+ tokens = reader.readlines()
152
+ for index, token in enumerate(tokens):
153
+ token = token.rstrip("\n")
154
+ vocab[token] = index
155
+ return vocab
156
+
157
+ @result_to_dataframe(["retrieved_contents", "retrieved_ids", "retrieve_scores"])
158
+ def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
159
+ queries, contents, _, ids = self.cast_to_run(previous_result)
160
+ top_k = kwargs.pop("top_k")
161
+ batch = kwargs.pop("batch", 64)
162
+ return self._pure(queries, contents, ids, top_k, batch)
163
+
164
+ def _pure(
165
+ self,
166
+ queries: List[str],
167
+ contents_list: List[List[str]],
168
+ ids_list: List[List[str]],
169
+ top_k: int,
170
+ batch: int = 64,
171
+ ) -> Tuple[List[List[str]], List[List[str]], List[List[float]]]:
172
+ """
173
+ Rerank a list of contents with FlashRank rerank models.
174
+
175
+ :param queries: The list of queries to use for reranking
176
+ :param contents_list: The list of lists of contents to rerank
177
+ :param ids_list: The list of lists of ids retrieved from the initial ranking
178
+ :param top_k: The number of passages to be retrieved
179
+ :param batch: The number of queries to be processed in a batch
180
+ :return: Tuple of lists containing the reranked contents, ids, and scores
181
+ """
182
+ nested_list = [
183
+ list(map(lambda x: [query, x], content_list))
184
+ for query, content_list in zip(queries, contents_list)
185
+ ]
186
+
187
+ rerank_scores = flatten_apply(
188
+ flashrank_run_model,
189
+ nested_list,
190
+ session=self.session,
191
+ batch_size=batch,
192
+ tokenizer=self.tokenizer,
193
+ )
194
+
195
+ df = pd.DataFrame(
196
+ {
197
+ "contents": contents_list,
198
+ "ids": ids_list,
199
+ "scores": rerank_scores,
200
+ }
201
+ )
202
+ df[["contents", "ids", "scores"]] = df.apply(
203
+ sort_by_scores, axis=1, result_type="expand"
204
+ )
205
+ results = select_top_k(df, ["contents", "ids", "scores"], top_k)
206
+
207
+ return (
208
+ results["contents"].tolist(),
209
+ results["ids"].tolist(),
210
+ results["scores"].tolist(),
211
+ )
212
+
213
+
214
+ def flashrank_run_model(input_texts, tokenizer, session, batch_size: int):
215
+ batch_input_texts = make_batch(input_texts, batch_size)
216
+ results = []
217
+
218
+ for batch_texts in tqdm(batch_input_texts):
219
+ input_text = tokenizer.encode_batch(batch_texts)
220
+ input_ids = np.array([e.ids for e in input_text])
221
+ token_type_ids = np.array([e.type_ids for e in input_text])
222
+ attention_mask = np.array([e.attention_mask for e in input_text])
223
+
224
+ use_token_type_ids = token_type_ids is not None and not np.all(
225
+ token_type_ids == 0
226
+ )
227
+
228
+ onnx_input = {
229
+ "input_ids": input_ids.astype(np.int64),
230
+ "attention_mask": attention_mask.astype(np.int64),
231
+ }
232
+ if use_token_type_ids:
233
+ onnx_input["token_type_ids"] = token_type_ids.astype(np.int64)
234
+
235
+ outputs = session.run(None, onnx_input)
236
+
237
+ logits = outputs[0]
238
+
239
+ if logits.shape[1] == 1:
240
+ scores = 1 / (1 + np.exp(-logits.flatten()))
241
+ else:
242
+ exp_logits = np.exp(logits)
243
+ scores = exp_logits[:, 1] / np.sum(exp_logits, axis=1)
244
+ results.extend(scores)
245
+ return results