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,239 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import uuid
|
|
3
|
+
from typing import Callable, Optional, List
|
|
4
|
+
|
|
5
|
+
import chromadb
|
|
6
|
+
import numpy as np
|
|
7
|
+
import pandas as pd
|
|
8
|
+
from tqdm import tqdm
|
|
9
|
+
|
|
10
|
+
import autorag
|
|
11
|
+
from autorag.nodes.semanticretrieval.vectordb import vectordb_ingest_api, vectordb_pure
|
|
12
|
+
from autorag.utils.util import (
|
|
13
|
+
save_parquet_safe,
|
|
14
|
+
fetch_contents,
|
|
15
|
+
get_event_loop,
|
|
16
|
+
process_batch,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
logger = logging.getLogger("AutoRAG")
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def make_single_content_qa(
|
|
23
|
+
corpus_df: pd.DataFrame,
|
|
24
|
+
content_size: int,
|
|
25
|
+
qa_creation_func: Callable,
|
|
26
|
+
output_filepath: Optional[str] = None,
|
|
27
|
+
upsert: bool = False,
|
|
28
|
+
random_state: int = 42,
|
|
29
|
+
cache_batch: int = 32,
|
|
30
|
+
**kwargs,
|
|
31
|
+
) -> pd.DataFrame:
|
|
32
|
+
"""
|
|
33
|
+
Make single content (single-hop, single-document) QA dataset using given qa_creation_func.
|
|
34
|
+
It generates a single content QA dataset, which means its retrieval ground truth will be only one.
|
|
35
|
+
It is the most basic form of QA dataset.
|
|
36
|
+
|
|
37
|
+
:param corpus_df: The corpus dataframe to make QA dataset from.
|
|
38
|
+
:param content_size: This function will generate QA dataset for the given number of contents.
|
|
39
|
+
:param qa_creation_func: The function to create QA pairs.
|
|
40
|
+
You can use like `generate_qa_llama_index` or `generate_qa_llama_index_by_ratio`.
|
|
41
|
+
The input func must have `contents` parameter for the list of content string.
|
|
42
|
+
:param output_filepath: Optional filepath to save the parquet file.
|
|
43
|
+
If None, the function will return the processed_data as pd.DataFrame, but do not save as parquet.
|
|
44
|
+
File directory must exist. File extension must be .parquet
|
|
45
|
+
:param upsert: If true, the function will overwrite the existing file if it exists.
|
|
46
|
+
Default is False.
|
|
47
|
+
:param random_state: The random state for sampling corpus from the given corpus_df.
|
|
48
|
+
:param cache_batch: The number of batches to use for caching the generated QA dataset.
|
|
49
|
+
When the cache_batch size data is generated, the dataset will save to the designated output_filepath.
|
|
50
|
+
If the cache_batch size is too small, the process time will be longer.
|
|
51
|
+
:param kwargs: The keyword arguments for qa_creation_func.
|
|
52
|
+
:return: QA dataset dataframe.
|
|
53
|
+
You can save this as parquet file to use at AutoRAG.
|
|
54
|
+
"""
|
|
55
|
+
assert content_size > 0, "content_size must be greater than 0."
|
|
56
|
+
if content_size > len(corpus_df):
|
|
57
|
+
logger.warning(
|
|
58
|
+
f"content_size {content_size} is larger than the corpus size {len(corpus_df)}. "
|
|
59
|
+
"Setting content_size to the corpus size."
|
|
60
|
+
)
|
|
61
|
+
content_size = len(corpus_df)
|
|
62
|
+
sampled_corpus = corpus_df.sample(n=content_size, random_state=random_state)
|
|
63
|
+
sampled_corpus = sampled_corpus.reset_index(drop=True)
|
|
64
|
+
|
|
65
|
+
def make_query_generation_gt(row):
|
|
66
|
+
return row["qa"]["query"], row["qa"]["generation_gt"]
|
|
67
|
+
|
|
68
|
+
qa_data = pd.DataFrame()
|
|
69
|
+
for idx, i in tqdm(enumerate(range(0, len(sampled_corpus), cache_batch))):
|
|
70
|
+
qa = qa_creation_func(
|
|
71
|
+
contents=sampled_corpus["contents"].tolist()[i : i + cache_batch], **kwargs
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
temp_qa_data = pd.DataFrame(
|
|
75
|
+
{
|
|
76
|
+
"qa": qa,
|
|
77
|
+
"retrieval_gt": sampled_corpus["doc_id"].tolist()[i : i + cache_batch],
|
|
78
|
+
}
|
|
79
|
+
)
|
|
80
|
+
temp_qa_data = temp_qa_data.explode("qa", ignore_index=True)
|
|
81
|
+
temp_qa_data["qid"] = [str(uuid.uuid4()) for _ in range(len(temp_qa_data))]
|
|
82
|
+
temp_qa_data[["query", "generation_gt"]] = temp_qa_data.apply(
|
|
83
|
+
make_query_generation_gt, axis=1, result_type="expand"
|
|
84
|
+
)
|
|
85
|
+
temp_qa_data = temp_qa_data.drop(columns=["qa"])
|
|
86
|
+
|
|
87
|
+
temp_qa_data["retrieval_gt"] = temp_qa_data["retrieval_gt"].apply(
|
|
88
|
+
lambda x: [[x]]
|
|
89
|
+
)
|
|
90
|
+
temp_qa_data["generation_gt"] = temp_qa_data["generation_gt"].apply(
|
|
91
|
+
lambda x: [x]
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
if idx == 0:
|
|
95
|
+
qa_data = temp_qa_data
|
|
96
|
+
else:
|
|
97
|
+
qa_data = pd.concat([qa_data, temp_qa_data], ignore_index=True)
|
|
98
|
+
if output_filepath is not None:
|
|
99
|
+
save_parquet_safe(qa_data, output_filepath, upsert=upsert)
|
|
100
|
+
|
|
101
|
+
return qa_data
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def make_qa_with_existing_qa(
|
|
105
|
+
corpus_df: pd.DataFrame,
|
|
106
|
+
existing_query_df: pd.DataFrame,
|
|
107
|
+
content_size: int,
|
|
108
|
+
answer_creation_func: Optional[Callable] = None,
|
|
109
|
+
exist_gen_gt: Optional[bool] = False,
|
|
110
|
+
output_filepath: Optional[str] = None,
|
|
111
|
+
embedding_model: str = "openai_embed_3_large",
|
|
112
|
+
collection: Optional[chromadb.Collection] = None,
|
|
113
|
+
upsert: bool = False,
|
|
114
|
+
random_state: int = 42,
|
|
115
|
+
cache_batch: int = 32,
|
|
116
|
+
top_k: int = 3,
|
|
117
|
+
**kwargs,
|
|
118
|
+
) -> pd.DataFrame:
|
|
119
|
+
"""
|
|
120
|
+
Make single-hop QA dataset using given qa_creation_func and existing queries.
|
|
121
|
+
|
|
122
|
+
:param corpus_df: The corpus dataframe to make QA dataset from.
|
|
123
|
+
:param existing_query_df: Dataframe containing existing queries to use for QA pair creation.
|
|
124
|
+
:param content_size: This function will generate QA dataset for the given number of contents.
|
|
125
|
+
:param answer_creation_func: Optional function to create answer with input query.
|
|
126
|
+
If exist_gen_gt is False, this function must be given.
|
|
127
|
+
:param exist_gen_gt: Optional boolean to use existing generation_gt.
|
|
128
|
+
If True, the existing_query_df must have 'generation_gt' column.
|
|
129
|
+
If False, the answer_creation_func must be given.
|
|
130
|
+
:param output_filepath: Optional filepath to save the parquet file.
|
|
131
|
+
:param embedding_model: The embedding model to use for vectorization.
|
|
132
|
+
You can add your own embedding model in the autorag.embedding_models.
|
|
133
|
+
Please refer to how to add an embedding model in this doc: https://marker-inc-korea.github.io/AutoRAG/local_model.html
|
|
134
|
+
The default is 'openai_embed_3_large'.
|
|
135
|
+
:param collection: The chromadb collection to use for vector DB.
|
|
136
|
+
You can make any chromadb collection and use it here.
|
|
137
|
+
If you already ingested the corpus_df to the collection, the embedding process will not be repeated.
|
|
138
|
+
The default is None. If None, it makes a temporary collection.
|
|
139
|
+
:param upsert: If true, the function will overwrite the existing file if it exists.
|
|
140
|
+
:param random_state: The random state for sampling corpus from the given corpus_df.
|
|
141
|
+
:param cache_batch: The number of batches to use for caching the generated QA dataset.
|
|
142
|
+
:param top_k: The number of sources to refer by model.
|
|
143
|
+
Default is 3.
|
|
144
|
+
:param kwargs: The keyword arguments for qa_creation_func.
|
|
145
|
+
:return: QA dataset dataframe.
|
|
146
|
+
"""
|
|
147
|
+
raise DeprecationWarning("This function is deprecated.")
|
|
148
|
+
assert "query" in existing_query_df.columns, (
|
|
149
|
+
"existing_query_df must have 'query' column."
|
|
150
|
+
)
|
|
151
|
+
|
|
152
|
+
if exist_gen_gt:
|
|
153
|
+
assert "generation_gt" in existing_query_df.columns, (
|
|
154
|
+
"existing_query_df must have 'generation_gt' column."
|
|
155
|
+
)
|
|
156
|
+
else:
|
|
157
|
+
assert answer_creation_func is not None, (
|
|
158
|
+
"answer_creation_func must be given when exist_gen_gt is False."
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
assert content_size > 0, "content_size must be greater than 0."
|
|
162
|
+
if content_size > len(corpus_df):
|
|
163
|
+
logger.warning(
|
|
164
|
+
f"content_size {content_size} is larger than the corpus size {len(corpus_df)}. "
|
|
165
|
+
"Setting content_size to the corpus size."
|
|
166
|
+
)
|
|
167
|
+
content_size = len(corpus_df)
|
|
168
|
+
|
|
169
|
+
logger.info("Loading local embedding model...")
|
|
170
|
+
embeddings = autorag.embedding_models[embedding_model]()
|
|
171
|
+
|
|
172
|
+
# Vector DB creation
|
|
173
|
+
if collection is None:
|
|
174
|
+
chroma_client = chromadb.Client()
|
|
175
|
+
collection_name = "auto-rag"
|
|
176
|
+
collection = chroma_client.get_or_create_collection(collection_name)
|
|
177
|
+
|
|
178
|
+
# embed corpus_df
|
|
179
|
+
vectordb_ingest_api(collection, corpus_df, embeddings)
|
|
180
|
+
query_embeddings = embeddings.get_text_embedding_batch(
|
|
181
|
+
existing_query_df["query"].tolist()
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
loop = get_event_loop()
|
|
185
|
+
tasks = [
|
|
186
|
+
vectordb_pure([query_embedding], top_k, collection)
|
|
187
|
+
for query_embedding in query_embeddings
|
|
188
|
+
]
|
|
189
|
+
results = loop.run_until_complete(process_batch(tasks, batch_size=cache_batch))
|
|
190
|
+
retrieved_ids = list(map(lambda x: x[0], results))
|
|
191
|
+
|
|
192
|
+
retrieved_contents: List[List[str]] = fetch_contents(corpus_df, retrieved_ids)
|
|
193
|
+
input_passage_strs: List[str] = list(
|
|
194
|
+
map(
|
|
195
|
+
lambda x: "\n".join(
|
|
196
|
+
[f"Document {i + 1}\n{content}" for i, content in enumerate(x)]
|
|
197
|
+
),
|
|
198
|
+
retrieved_contents,
|
|
199
|
+
)
|
|
200
|
+
)
|
|
201
|
+
|
|
202
|
+
retrieved_qa_df = pd.DataFrame(
|
|
203
|
+
{
|
|
204
|
+
"qid": [str(uuid.uuid4()) for _ in range(len(existing_query_df))],
|
|
205
|
+
"query": existing_query_df["query"].tolist(),
|
|
206
|
+
"retrieval_gt": list(map(lambda x: [x], retrieved_ids)),
|
|
207
|
+
"input_passage_str": input_passage_strs,
|
|
208
|
+
}
|
|
209
|
+
)
|
|
210
|
+
|
|
211
|
+
if exist_gen_gt:
|
|
212
|
+
generation_gt = existing_query_df["generation_gt"].tolist()
|
|
213
|
+
if isinstance(generation_gt[0], np.ndarray):
|
|
214
|
+
retrieved_qa_df["generation_gt"] = generation_gt
|
|
215
|
+
else:
|
|
216
|
+
raise ValueError(
|
|
217
|
+
"In existing_query_df, generation_gt (per query) must be in the form of List[str]."
|
|
218
|
+
)
|
|
219
|
+
|
|
220
|
+
sample_qa_df = retrieved_qa_df.sample(
|
|
221
|
+
n=min(content_size, len(retrieved_qa_df)), random_state=random_state
|
|
222
|
+
)
|
|
223
|
+
|
|
224
|
+
qa_df = sample_qa_df.copy(deep=True)
|
|
225
|
+
qa_df.drop(columns=["input_passage_str"], inplace=True)
|
|
226
|
+
|
|
227
|
+
if not exist_gen_gt:
|
|
228
|
+
generation_gt = answer_creation_func(
|
|
229
|
+
contents=sample_qa_df["input_passage_str"].tolist(),
|
|
230
|
+
queries=sample_qa_df["query"].tolist(),
|
|
231
|
+
batch=cache_batch,
|
|
232
|
+
**kwargs,
|
|
233
|
+
)
|
|
234
|
+
qa_df["generation_gt"] = generation_gt
|
|
235
|
+
|
|
236
|
+
if output_filepath is not None:
|
|
237
|
+
save_parquet_safe(qa_df, output_filepath, upsert=upsert)
|
|
238
|
+
|
|
239
|
+
return qa_df
|
|
@@ -0,0 +1,253 @@
|
|
|
1
|
+
import os.path
|
|
2
|
+
import random
|
|
3
|
+
from typing import Optional, List, Dict, Any
|
|
4
|
+
|
|
5
|
+
import pandas as pd
|
|
6
|
+
from llama_index.core.base.llms.types import ChatMessage, MessageRole
|
|
7
|
+
from llama_index.core.llms import LLM
|
|
8
|
+
|
|
9
|
+
from autorag.utils.util import process_batch, get_event_loop
|
|
10
|
+
|
|
11
|
+
package_dir = os.path.dirname(os.path.realpath(__file__))
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def generate_qa_llama_index(
|
|
15
|
+
llm: LLM,
|
|
16
|
+
contents: List[str],
|
|
17
|
+
prompt: Optional[str] = None,
|
|
18
|
+
question_num_per_content: int = 1,
|
|
19
|
+
max_retries: int = 3,
|
|
20
|
+
batch: int = 4,
|
|
21
|
+
) -> List[List[Dict]]:
|
|
22
|
+
"""
|
|
23
|
+
Generate a qa set from the list of contents.
|
|
24
|
+
It uses a single prompt for all contents.
|
|
25
|
+
If you want to use more than one prompt for generating qa,
|
|
26
|
+
you can consider using generate_qa_llama_index_by_ratio.
|
|
27
|
+
|
|
28
|
+
:param llm: Llama index model
|
|
29
|
+
:param contents: List of content strings.
|
|
30
|
+
:param prompt: The prompt to use for the qa generation.
|
|
31
|
+
The prompt must include the following placeholders:
|
|
32
|
+
- {{text}}: The content string
|
|
33
|
+
- {{num_questions}}: The number of questions to generate
|
|
34
|
+
As default, the prompt is set to the default prompt for the question type.
|
|
35
|
+
:param question_num_per_content: Number of questions to generate for each content.
|
|
36
|
+
Default is 1.
|
|
37
|
+
:param max_retries: The maximum number of retries when generated question number is not equal to the target number.
|
|
38
|
+
Default is 3.
|
|
39
|
+
:param batch: The batch size to process asynchronously.
|
|
40
|
+
Default is 4.
|
|
41
|
+
:return: 2-d list of dictionaries containing the query and generation_gt.
|
|
42
|
+
"""
|
|
43
|
+
# load default prompt
|
|
44
|
+
if prompt is None:
|
|
45
|
+
prompt = open(
|
|
46
|
+
os.path.join(package_dir, "llama_index_default_prompt.txt"), "r"
|
|
47
|
+
).read()
|
|
48
|
+
|
|
49
|
+
tasks = [
|
|
50
|
+
async_qa_gen_llama_index(
|
|
51
|
+
content, llm, prompt, question_num_per_content, max_retries
|
|
52
|
+
)
|
|
53
|
+
for content in contents
|
|
54
|
+
]
|
|
55
|
+
loops = get_event_loop()
|
|
56
|
+
results = loops.run_until_complete(process_batch(tasks, batch))
|
|
57
|
+
return results
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def generate_answers(
|
|
61
|
+
llm: LLM,
|
|
62
|
+
contents: List[str],
|
|
63
|
+
queries: List[str],
|
|
64
|
+
batch: int = 4,
|
|
65
|
+
) -> List[List[Dict]]:
|
|
66
|
+
"""
|
|
67
|
+
Generate qa sets from the list of contents using existing queries.
|
|
68
|
+
|
|
69
|
+
:param llm: Llama index model
|
|
70
|
+
:param contents: List of content strings.
|
|
71
|
+
:param queries: List of existing queries.
|
|
72
|
+
:param batch: The batch size to process asynchronously.
|
|
73
|
+
:return: 2-d list of dictionaries containing the query and generation_gt.
|
|
74
|
+
"""
|
|
75
|
+
|
|
76
|
+
tasks = [
|
|
77
|
+
generate_basic_answer(llm, content, query)
|
|
78
|
+
for content, query in zip(contents, queries)
|
|
79
|
+
]
|
|
80
|
+
loops = get_event_loop()
|
|
81
|
+
results = loops.run_until_complete(process_batch(tasks, batch))
|
|
82
|
+
return results
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def generate_qa_llama_index_by_ratio(
|
|
86
|
+
llm: LLM,
|
|
87
|
+
contents: List[str],
|
|
88
|
+
prompts_ratio: Dict,
|
|
89
|
+
question_num_per_content: int = 1,
|
|
90
|
+
max_retries: int = 3,
|
|
91
|
+
random_state: int = 42,
|
|
92
|
+
batch: int = 4,
|
|
93
|
+
) -> List[List[Dict]]:
|
|
94
|
+
"""
|
|
95
|
+
Generate a qa set from the list of contents.
|
|
96
|
+
You can set the ratio of prompts that you want to use for generating qa.
|
|
97
|
+
It distributes the number of questions to generate for each content by the ratio randomly.
|
|
98
|
+
|
|
99
|
+
:param llm: Llama index model
|
|
100
|
+
:param contents: List of content strings.
|
|
101
|
+
:param prompts_ratio: Dictionary of prompt paths and their ratios.
|
|
102
|
+
Example: {"prompt/prompt1.txt": 0.5, "prompt/prompt2.txt": 0.5}
|
|
103
|
+
The value sum doesn't have to be 1.
|
|
104
|
+
The path must be the absolute path, and the file must exist.
|
|
105
|
+
Plus, it has to be a text file which contains proper prompt.
|
|
106
|
+
Each prompt must contain the following placeholders:
|
|
107
|
+
- {{text}}: The content string
|
|
108
|
+
- {{num_questions}}: The number of questions to generate
|
|
109
|
+
:param question_num_per_content: Number of questions to generate for each content.
|
|
110
|
+
Default is 1.
|
|
111
|
+
:param max_retries: The maximum number of retries when generated question number is not equal to the target number.
|
|
112
|
+
Default is 3.
|
|
113
|
+
:param random_state: Random seed
|
|
114
|
+
Default is 42.
|
|
115
|
+
:param batch: The batch size to process asynchronously.
|
|
116
|
+
Default is 4.
|
|
117
|
+
:return: 2-d list of dictionaries containing the query and generation_gt.
|
|
118
|
+
"""
|
|
119
|
+
prompts = list(map(lambda path: open(path, "r").read(), prompts_ratio.keys()))
|
|
120
|
+
assert all([validate_llama_index_prompt(prompt) for prompt in prompts])
|
|
121
|
+
|
|
122
|
+
content_indices = list(range(len(contents)))
|
|
123
|
+
random.seed(random_state)
|
|
124
|
+
random.shuffle(content_indices)
|
|
125
|
+
|
|
126
|
+
slice_content_indices: List[List[str]] = distribute_list_by_ratio(
|
|
127
|
+
content_indices, list(prompts_ratio.values())
|
|
128
|
+
)
|
|
129
|
+
temp_df = pd.DataFrame({"idx": slice_content_indices, "prompt": prompts})
|
|
130
|
+
temp_df = temp_df.explode("idx", ignore_index=True)
|
|
131
|
+
temp_df = temp_df.sort_values(by="idx", ascending=True)
|
|
132
|
+
|
|
133
|
+
final_df = pd.DataFrame({"content": contents, "prompt": temp_df["prompt"].tolist()})
|
|
134
|
+
|
|
135
|
+
tasks = [
|
|
136
|
+
async_qa_gen_llama_index(
|
|
137
|
+
content, llm, prompt, question_num_per_content, max_retries
|
|
138
|
+
)
|
|
139
|
+
for content, prompt in zip(
|
|
140
|
+
final_df["content"].tolist(), final_df["prompt"].tolist()
|
|
141
|
+
)
|
|
142
|
+
]
|
|
143
|
+
|
|
144
|
+
loops = get_event_loop()
|
|
145
|
+
results = loops.run_until_complete(process_batch(tasks, batch))
|
|
146
|
+
|
|
147
|
+
return results
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
async def async_qa_gen_llama_index(
|
|
151
|
+
content: str,
|
|
152
|
+
llm: LLM,
|
|
153
|
+
prompt: str,
|
|
154
|
+
question_num: int = 1,
|
|
155
|
+
max_retries: int = 3,
|
|
156
|
+
):
|
|
157
|
+
"""
|
|
158
|
+
Generate a qa set by using the given content and the llama index model.
|
|
159
|
+
You must select the question type.
|
|
160
|
+
|
|
161
|
+
:param content: Content string
|
|
162
|
+
:param llm: Llama index model
|
|
163
|
+
:param prompt: The prompt to use for the qa generation.
|
|
164
|
+
The prompt must include the following placeholders:
|
|
165
|
+
- {{text}}: The content string
|
|
166
|
+
- {{num_questions}}: The number of questions to generate
|
|
167
|
+
:param question_num: The number of questions to generate
|
|
168
|
+
:param max_retries: Maximum number of retries when generated question number is not equal to the target number
|
|
169
|
+
:return: List of dictionaries containing the query and generation_gt
|
|
170
|
+
"""
|
|
171
|
+
validate_llama_index_prompt(prompt)
|
|
172
|
+
|
|
173
|
+
async def generate(content: str, llm: LLM):
|
|
174
|
+
for _ in range(max_retries):
|
|
175
|
+
output = await llm.acomplete(
|
|
176
|
+
prompt.replace("{{text}}", content).replace(
|
|
177
|
+
"{{num_questions}}", str(question_num)
|
|
178
|
+
)
|
|
179
|
+
)
|
|
180
|
+
result = parse_output(output.text)
|
|
181
|
+
if len(result) == question_num:
|
|
182
|
+
return result
|
|
183
|
+
raise InterruptedError(
|
|
184
|
+
f"Failed to generate output of length {question_num} after {max_retries} retries."
|
|
185
|
+
)
|
|
186
|
+
|
|
187
|
+
return await generate(content, llm)
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
async def generate_basic_answer(llm: LLM, passage_str: str, query: str) -> str:
|
|
191
|
+
basic_answer_system_prompt = """You are an AI assistant to answer the given question in the provide evidence text.
|
|
192
|
+
You can find the evidence from the given text about question, and you have to write a proper answer to the given question.
|
|
193
|
+
You have to preserve the question's language at the answer.
|
|
194
|
+
For example, if the input question is Korean, the output answer must be in Korean.
|
|
195
|
+
"""
|
|
196
|
+
user_prompt = f"Text:\n<|text_start|>\n{passage_str}\n<|text_end|>\n\nQuestion:\n{query}\n\nAnswer:"
|
|
197
|
+
|
|
198
|
+
response = await llm.achat(
|
|
199
|
+
messages=[
|
|
200
|
+
ChatMessage(role=MessageRole.SYSTEM, content=basic_answer_system_prompt),
|
|
201
|
+
ChatMessage(role=MessageRole.USER, content=user_prompt),
|
|
202
|
+
],
|
|
203
|
+
temperature=1.0,
|
|
204
|
+
)
|
|
205
|
+
return response.message.content
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def validate_llama_index_prompt(prompt: str) -> bool:
|
|
209
|
+
"""
|
|
210
|
+
Validate the prompt for the llama index model.
|
|
211
|
+
The prompt must include the following placeholders:
|
|
212
|
+
- {{text}}: The content string
|
|
213
|
+
- {{num_questions}}: The number of questions to generate
|
|
214
|
+
"""
|
|
215
|
+
if "{{text}}" not in prompt:
|
|
216
|
+
raise ValueError("The prompt must include the placeholder {{text}}.")
|
|
217
|
+
if "{{num_questions}}" not in prompt:
|
|
218
|
+
raise ValueError("The prompt must include the placeholder {{num_questions}}.")
|
|
219
|
+
return True
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def parse_output(result: str) -> List[Dict]:
|
|
223
|
+
result = result.strip()
|
|
224
|
+
result = result.split("[Q]:")
|
|
225
|
+
final_result = list()
|
|
226
|
+
for res in result:
|
|
227
|
+
res = res.strip()
|
|
228
|
+
if res and "\n[A]:" in res:
|
|
229
|
+
qa = res.split("\n[A]:")
|
|
230
|
+
final_result.append(
|
|
231
|
+
{"query": qa[0].strip(), "generation_gt": qa[1].strip()}
|
|
232
|
+
)
|
|
233
|
+
return final_result
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def distribute_list_by_ratio(input_list, ratio) -> List[List[Any]]:
|
|
237
|
+
total_ratio = sum(ratio)
|
|
238
|
+
total_length = len(input_list)
|
|
239
|
+
|
|
240
|
+
# Calculate the length of each slice
|
|
241
|
+
slice_lengths = [int((r / total_ratio) * total_length) for r in ratio]
|
|
242
|
+
|
|
243
|
+
# Adjust the last slice in case of rounding issues
|
|
244
|
+
slice_lengths[-1] = total_length - sum(slice_lengths[:-1])
|
|
245
|
+
|
|
246
|
+
slices = []
|
|
247
|
+
start = 0
|
|
248
|
+
for length in slice_lengths:
|
|
249
|
+
end = start + length
|
|
250
|
+
slices.append(input_list[start:end])
|
|
251
|
+
start = end
|
|
252
|
+
|
|
253
|
+
return slices
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
You're an AI tasked to convert Text into a question and answer set.
|
|
2
|
+
Cover as many details from Text as possible in the QnA set.
|
|
3
|
+
|
|
4
|
+
Instructions:
|
|
5
|
+
1. Both Questions and Answers MUST BE extracted from given Text
|
|
6
|
+
2. Answers must be full sentences
|
|
7
|
+
3. Questions should be as detailed as possible from Text
|
|
8
|
+
4. Output must always have the provided number of QnAs
|
|
9
|
+
5. Create questions that ask about information from the Text
|
|
10
|
+
6. MUST include specific keywords from the Text.
|
|
11
|
+
7. Do not mention any of these in the questions: "in the given text", "in the provided information", etc.
|
|
12
|
+
|
|
13
|
+
Question examples:
|
|
14
|
+
1. How do owen and riggs know each other?
|
|
15
|
+
2. What does the word fore "mean" in golf?
|
|
16
|
+
3. What makes charging bull in nyc popular to tourists?
|
|
17
|
+
4. What kind of pistol does the army use?
|
|
18
|
+
5. Who was the greatest violin virtuoso in the romantic period?
|
|
19
|
+
<|separator|>
|
|
20
|
+
|
|
21
|
+
Text:
|
|
22
|
+
<|text_start|>
|
|
23
|
+
Mark Hamill as Luke Skywalker : One of the last living Jedi , trained by Obi - Wan and Yoda , who is also a skilled X-wing fighter pilot allied with the Rebellion .
|
|
24
|
+
Harrison Ford as Han Solo : A rogue smuggler , who aids the Rebellion against the Empire . Han is Luke and Leia 's friend , as well as Leia 's love interest .
|
|
25
|
+
Carrie Fisher as Leia Organa : The former Princess of the destroyed planet Alderaan , who joins the Rebellion ; Luke 's twin sister , and Han 's love interest .
|
|
26
|
+
Billy Dee Williams as Lando Calrissian : The former Baron Administrator of Cloud City and one of Han 's friends who aids the Rebellion .
|
|
27
|
+
Anthony Daniels as C - 3PO : A humanoid protocol droid , who sides with the Rebellion .
|
|
28
|
+
Peter Mayhew as Chewbacca : A Wookiee who is Han 's longtime friend , who takes part in the Rebellion .
|
|
29
|
+
Kenny Baker as R2 - D2 : An astromech droid , bought by Luke ; and long - time friend to C - 3PO . He also portrays a GONK power droid in the background .
|
|
30
|
+
Ian McDiarmid as the Emperor : The evil founding supreme ruler of the Galactic Empire , and Vader 's Sith Master .
|
|
31
|
+
Frank Oz as Yoda : The wise , centuries - old Grand Master of the Jedi , who is Luke 's self - exiled Jedi Master living on Dagobah . After dying , he reappears to Luke as a Force - ghost . Yoda 's Puppetry was assisted by Mike Quinn .
|
|
32
|
+
David Prowse as Darth Vader / Anakin Skywalker : A powerful Sith lord and the second in command of the Galactic Empire ; Luke and Leia 's father .
|
|
33
|
+
<|text_end|>
|
|
34
|
+
Output with 4 QnAs:
|
|
35
|
+
<|separator|>
|
|
36
|
+
|
|
37
|
+
[Q]: who played luke father in return of the jedi
|
|
38
|
+
[A]: David Prowse acted as Darth Vader, a.k.a Anakin Skywalker, which is Luke and Leia's father.
|
|
39
|
+
[Q]: Who is Han Solo's best friend? And what species is he?
|
|
40
|
+
[A]: Han Solo's best friend is Chewbacca, who is a Wookiee.
|
|
41
|
+
[Q]: Who played luke's teacher in the return of the jedi
|
|
42
|
+
[A]: Yoda, the wise, centuries-old Grand Master of the Jedi, who is Luke's self-exiled Jedi Master living on Dagobah, was played by Frank Oz.
|
|
43
|
+
Also, there is a mention of Obi-Wan Kenobi, who trained Luke Skywalker.
|
|
44
|
+
But I can't find who played Obi-Wan Kenobi in the given text.
|
|
45
|
+
[Q]: Where Yoda lives in the return of the jedi?
|
|
46
|
+
[A]: Yoda, the Jedi Master, lives on Dagobah.
|
|
47
|
+
<|separator|>
|
|
48
|
+
|
|
49
|
+
Text:
|
|
50
|
+
<|text_start|>
|
|
51
|
+
{{text}}
|
|
52
|
+
<|text_end|>
|
|
53
|
+
Output with {{num_questions}} QnAs:
|
|
54
|
+
<|separator|>
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
import uuid
|
|
2
|
+
from typing import Optional
|
|
3
|
+
|
|
4
|
+
import pandas as pd
|
|
5
|
+
from langchain_core.embeddings import Embeddings
|
|
6
|
+
from langchain_core.language_models import BaseChatModel
|
|
7
|
+
from langchain_openai import ChatOpenAI, OpenAIEmbeddings
|
|
8
|
+
|
|
9
|
+
from autorag.data.utils.util import corpus_df_to_langchain_documents
|
|
10
|
+
from autorag.utils import cast_qa_dataset
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def generate_qa_ragas(
|
|
14
|
+
corpus_df: pd.DataFrame,
|
|
15
|
+
test_size: int,
|
|
16
|
+
distributions: Optional[dict] = None,
|
|
17
|
+
generator_llm: Optional[BaseChatModel] = None,
|
|
18
|
+
critic_llm: Optional[BaseChatModel] = None,
|
|
19
|
+
embedding_model: Optional[Embeddings] = None,
|
|
20
|
+
**kwargs,
|
|
21
|
+
) -> pd.DataFrame:
|
|
22
|
+
"""
|
|
23
|
+
QA dataset generation using RAGAS.
|
|
24
|
+
Returns qa dataset dataframe.
|
|
25
|
+
|
|
26
|
+
:param corpus_df: Corpus dataframe.
|
|
27
|
+
:param test_size: Number of queries to generate.
|
|
28
|
+
:param distributions: Distributions of different types of questions.
|
|
29
|
+
Default is "simple is 0.5, multi_context is 0.4, and reasoning is 0.1."
|
|
30
|
+
Each type of questions refers to Ragas evolution types.
|
|
31
|
+
:param generator_llm: Generator language model from Langchain.
|
|
32
|
+
:param critic_llm: Critic language model from Langchain.
|
|
33
|
+
:param embedding_model: Embedding model from Langchain.
|
|
34
|
+
:param kwargs: The additional option to pass to the 'generate_with_langchain_docs' method.
|
|
35
|
+
You can input 'with_debugging_logs', 'is_async', 'raise_exceptions', and 'run_config'.
|
|
36
|
+
:return: QA dataset dataframe.
|
|
37
|
+
"""
|
|
38
|
+
from ragas.testset import TestsetGenerator
|
|
39
|
+
from ragas.testset.evolutions import simple, reasoning, multi_context
|
|
40
|
+
|
|
41
|
+
if generator_llm is None:
|
|
42
|
+
generator_llm = ChatOpenAI(model="gpt-3.5-turbo-16k")
|
|
43
|
+
if critic_llm is None:
|
|
44
|
+
critic_llm = ChatOpenAI(model="gpt-4-turbo")
|
|
45
|
+
if embedding_model is None:
|
|
46
|
+
embedding_model = OpenAIEmbeddings()
|
|
47
|
+
if distributions is None:
|
|
48
|
+
distributions = {simple: 0.5, multi_context: 0.4, reasoning: 0.1}
|
|
49
|
+
|
|
50
|
+
assert sum(list(distributions.values())) == 1.0, "Sum of distributions must be 1.0"
|
|
51
|
+
|
|
52
|
+
generator = TestsetGenerator.from_langchain(
|
|
53
|
+
generator_llm, critic_llm, embedding_model
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
langchain_docs = corpus_df_to_langchain_documents(corpus_df)
|
|
57
|
+
|
|
58
|
+
test_df = generator.generate_with_langchain_docs(
|
|
59
|
+
langchain_docs, test_size, distributions=distributions, **kwargs
|
|
60
|
+
).to_pandas()
|
|
61
|
+
|
|
62
|
+
result_df = pd.DataFrame(
|
|
63
|
+
{
|
|
64
|
+
"qid": [str(uuid.uuid4()) for _ in range(len(test_df))],
|
|
65
|
+
"query": test_df["question"].tolist(),
|
|
66
|
+
"generation_gt": list(map(lambda x: x, test_df["ground_truth"].tolist())),
|
|
67
|
+
}
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
result_df["retrieval_gt"] = test_df["metadata"].apply(
|
|
71
|
+
lambda x: list(map(lambda y: y["filename"], x))
|
|
72
|
+
)
|
|
73
|
+
result_df = cast_qa_dataset(result_df)
|
|
74
|
+
|
|
75
|
+
return result_df
|