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,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
|