mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-05 15:20:30 +08:00
refa: simplify RAPTOR tree clustering and configuration (#17614)
This commit is contained in:
@@ -944,8 +944,6 @@ async def run_tree_templates(
|
||||
raptor_config=raptor_config,
|
||||
chat_mdl=chat_mdl_by_tid[template_id],
|
||||
embd_mdl=embedding_model,
|
||||
tree_builder="raptor",
|
||||
clustering_method="ahc",
|
||||
max_errors=3,
|
||||
)
|
||||
except Exception:
|
||||
|
||||
@@ -41,8 +41,6 @@ from rag.nlp import rag_tokenizer, search
|
||||
from rag.utils.raptor_utils import (
|
||||
collect_raptor_chunk_ids,
|
||||
collect_raptor_methods,
|
||||
get_raptor_clustering_method,
|
||||
get_raptor_tree_builder,
|
||||
get_skip_reason,
|
||||
make_raptor_summary_chunk_id,
|
||||
should_skip_raptor,
|
||||
@@ -118,8 +116,6 @@ class RaptorService:
|
||||
Tuple of (chunks, token_count, cleanup_raptor_chunks).
|
||||
"""
|
||||
raptor_config = kb_parser_config.get("raptor", {})
|
||||
tree_builder = get_raptor_tree_builder(raptor_config)
|
||||
clustering_method = get_raptor_clustering_method(raptor_config)
|
||||
vctr_nm = "q_%d_vec" % vector_size
|
||||
|
||||
res = []
|
||||
@@ -132,13 +128,9 @@ class RaptorService:
|
||||
|
||||
# Determine scope
|
||||
if raptor_config.get("scope", "file") == "file":
|
||||
res, tk_count = await self._run_file_level_raptor(
|
||||
raptor_config, tree_builder, clustering_method, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks
|
||||
)
|
||||
res, tk_count = await self._run_file_level_raptor(raptor_config, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks)
|
||||
else:
|
||||
res, tk_count = await self._run_dataset_level_raptor(
|
||||
raptor_config, tree_builder, clustering_method, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks
|
||||
)
|
||||
res, tk_count = await self._run_dataset_level_raptor(raptor_config, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks)
|
||||
|
||||
return res, tk_count, cleanup_raptor_chunks
|
||||
|
||||
@@ -158,7 +150,8 @@ class RaptorService:
|
||||
}
|
||||
return doc_info_by_id
|
||||
|
||||
async def _run_file_level_raptor(self, raptor_config, tree_builder, clustering_method, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks):
|
||||
async def _run_file_level_raptor(self, raptor_config, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks):
|
||||
tree_builder = "raptor"
|
||||
"""Run RAPTOR at file level (per document)."""
|
||||
ctx = self._task_context
|
||||
fake_doc_id = GRAPH_RAPTOR_FAKE_DOC_ID
|
||||
@@ -197,7 +190,7 @@ class RaptorService:
|
||||
continue
|
||||
|
||||
before_generate = len(res)
|
||||
new_chunks, new_tk_count = await self._generate_raptor(chunks, doc_id, raptor_config, chat_mdl, embd_mdl, tree_builder, clustering_method, max_errors, doc_info_by_id)
|
||||
new_chunks, new_tk_count = await self._generate_raptor(chunks, doc_id, raptor_config, chat_mdl, embd_mdl, max_errors, doc_info_by_id)
|
||||
res.extend(new_chunks)
|
||||
tk_count += new_tk_count
|
||||
|
||||
@@ -215,7 +208,8 @@ class RaptorService:
|
||||
|
||||
return res, tk_count
|
||||
|
||||
async def _run_dataset_level_raptor(self, raptor_config, tree_builder, clustering_method, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks):
|
||||
async def _run_dataset_level_raptor(self, raptor_config, chat_mdl, embd_mdl, vctr_nm, doc_ids, doc_info_by_id, max_errors, res, tk_count, cleanup_raptor_chunks):
|
||||
tree_builder = "raptor"
|
||||
"""Run RAPTOR at dataset level (all documents combined)."""
|
||||
ctx = self._task_context
|
||||
fake_doc_id = GRAPH_RAPTOR_FAKE_DOC_ID
|
||||
@@ -264,7 +258,7 @@ class RaptorService:
|
||||
return res, tk_count
|
||||
|
||||
before_generate = len(res)
|
||||
new_chunks, new_tk_count = await self._generate_raptor(chunks, fake_doc_id, raptor_config, chat_mdl, embd_mdl, tree_builder, clustering_method, max_errors, doc_info_by_id)
|
||||
new_chunks, new_tk_count = await self._generate_raptor(chunks, fake_doc_id, raptor_config, chat_mdl, embd_mdl, max_errors, doc_info_by_id)
|
||||
res.extend(new_chunks)
|
||||
tk_count += new_tk_count
|
||||
|
||||
@@ -354,8 +348,6 @@ class RaptorService:
|
||||
raptor_config: Dict,
|
||||
chat_mdl,
|
||||
embd_mdl,
|
||||
tree_builder: str,
|
||||
clustering_method: str,
|
||||
max_errors: int,
|
||||
doc_info_by_id: Dict,
|
||||
is_tree: bool = False,
|
||||
@@ -372,7 +364,6 @@ class RaptorService:
|
||||
ctx = self._task_context
|
||||
from rag.advanced_rag.knowlege_compile.raptor import RecursiveAbstractiveProcessing4TreeOrganizedRetrieval as Raptor
|
||||
|
||||
raptor_ext_config = raptor_config.get("ext") or {}
|
||||
assert chunks, "_generate_raptor must not be called with empty chunks"
|
||||
vctr_nm = "q_%d_vec" % len(chunks[0][1])
|
||||
|
||||
@@ -382,12 +373,9 @@ class RaptorService:
|
||||
embd_mdl,
|
||||
raptor_config["prompt"],
|
||||
raptor_config["max_token"],
|
||||
raptor_config["threshold"],
|
||||
max_errors=max_errors,
|
||||
tree_builder=tree_builder,
|
||||
clustering_method=clustering_method,
|
||||
psi_exact_max_leaves=raptor_ext_config.get("psi_exact_max_leaves", 4096),
|
||||
psi_bucket_size=raptor_ext_config.get("psi_bucket_size", 1024),
|
||||
clustering_threshold=float(raptor_config.get("clustering_threshold", 0.3)),
|
||||
clustering_ratio=float(raptor_config.get("clustering_ratio", 0.5)),
|
||||
)
|
||||
|
||||
# Seed each leaf with its own id as the start of its
|
||||
@@ -420,7 +408,6 @@ class RaptorService:
|
||||
raptor_config,
|
||||
doc_id,
|
||||
effective_doc_name,
|
||||
tree_builder,
|
||||
vctr_nm,
|
||||
)
|
||||
|
||||
@@ -432,7 +419,7 @@ class RaptorService:
|
||||
"docnm_kwd": effective_doc_name,
|
||||
"title_tks": rag_tokenizer.tokenize(effective_doc_name),
|
||||
"raptor_kwd": "raptor",
|
||||
"extra": {"raptor_method": tree_builder},
|
||||
"extra": {"raptor_method": "raptor"},
|
||||
"create_time": str(datetime.now()).replace("T", " ")[:19],
|
||||
"create_timestamp_flt": datetime.now().timestamp(),
|
||||
}
|
||||
@@ -465,7 +452,7 @@ class RaptorService:
|
||||
return res, tk_count
|
||||
|
||||
row_id = xxhash.xxh64(
|
||||
f"raptor_tree:{doc_id}:{tree_builder}".encode("utf-8", "surrogatepass"),
|
||||
f"raptor_tree:{doc_id}:raptor".encode("utf-8", "surrogatepass"),
|
||||
).hexdigest()
|
||||
row = {
|
||||
**doc,
|
||||
@@ -482,8 +469,6 @@ class RaptorService:
|
||||
raptor_config: Dict,
|
||||
chat_mdl,
|
||||
embd_mdl,
|
||||
tree_builder: str,
|
||||
clustering_method: str,
|
||||
max_errors: int,
|
||||
) -> Optional[Dict]:
|
||||
"""Build a RAPTOR tree dict for one document — no ES IO.
|
||||
@@ -497,19 +482,15 @@ class RaptorService:
|
||||
return None
|
||||
from rag.advanced_rag.knowlege_compile.raptor import RecursiveAbstractiveProcessing4TreeOrganizedRetrieval as Raptor
|
||||
|
||||
raptor_ext_config = raptor_config.get("ext") or {}
|
||||
raptor = Raptor(
|
||||
raptor_config.get("max_cluster", 64),
|
||||
chat_mdl,
|
||||
embd_mdl,
|
||||
raptor_config["prompt"],
|
||||
raptor_config["max_token"],
|
||||
raptor_config["threshold"],
|
||||
max_errors=max_errors,
|
||||
tree_builder=tree_builder,
|
||||
clustering_method=clustering_method,
|
||||
psi_exact_max_leaves=raptor_ext_config.get("psi_exact_max_leaves", 4096),
|
||||
psi_bucket_size=raptor_ext_config.get("psi_bucket_size", 1024),
|
||||
clustering_threshold=float(raptor_config.get("clustering_threshold", 0.3)),
|
||||
clustering_ratio=float(raptor_config.get("clustering_ratio", 0.5)),
|
||||
)
|
||||
|
||||
raptor_input = [(content, vctr, [chunk_id] if chunk_id else []) for content, vctr, chunk_id in chunks]
|
||||
@@ -537,7 +518,6 @@ class RaptorService:
|
||||
raptor_config,
|
||||
doc_id,
|
||||
effective_doc_name,
|
||||
tree_builder,
|
||||
vctr_nm,
|
||||
) -> Tuple[List[Dict], int]:
|
||||
"""Legacy per-summary materialization, kept only for PSI builds.
|
||||
@@ -563,7 +543,7 @@ class RaptorService:
|
||||
"docnm_kwd": effective_doc_name,
|
||||
"title_tks": rag_tokenizer.tokenize(effective_doc_name),
|
||||
"raptor_kwd": "raptor",
|
||||
"extra": {"raptor_method": tree_builder},
|
||||
"extra": {"raptor_method": "raptor"},
|
||||
}
|
||||
if ctx.pagerank:
|
||||
doc[PAGERANK_FLD] = int(ctx.pagerank)
|
||||
|
||||
@@ -395,14 +395,13 @@ class TaskHandler:
|
||||
{
|
||||
"raptor": {
|
||||
"use_raptor": True,
|
||||
"prompt": "Please summarize the following paragraphs. Be careful with the numbers, do not make things up. Paragraphs as following:\n {cluster_content}\nThe above is the content you need to summarize.",
|
||||
"max_token": 256,
|
||||
"threshold": 0.1,
|
||||
"prompt": "Summarize the paragraphs below without inventing facts or changing numbers.\nOutput exactly two parts in the same language as the source:\n1. First line: a concise title only.\n2. Following lines: a concise summary of the content.\nDo not output labels, Markdown headings, bullet points, or any other commentary.\n\nParagraphs:\n{cluster_content}",
|
||||
"max_token": 512,
|
||||
"clustering_threshold": 0.3,
|
||||
"clustering_ratio": 0.5,
|
||||
"max_cluster": 64,
|
||||
"random_seed": 0,
|
||||
"scope": "file",
|
||||
"clustering_method": "gmm",
|
||||
"tree_builder": "raptor",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user