feat: 从 main 分支迁移解析器与检索优化(不影响 API 接口)
- 新增 heading_rules.py 标题规则引擎,替换硬编码正则链 - 重写 mineru_parser.py:v2 格式兼容、表单自动检测 - manager.py: 跨页表格合并规则收紧(防误合并)+ LLM token 上限提升 - engine.py: embedding 缓存 + 同 section 表格邻居扩展 + prompt 改进 - router.py: 路由分类 token 上限提升至 512
This commit is contained in:
122
core/engine.py
122
core/engine.py
@@ -625,7 +625,7 @@ class RAGEngine:
|
|||||||
elif len(conditions) > 1:
|
elif len(conditions) > 1:
|
||||||
where_filter = {"$and": conditions}
|
where_filter = {"$and": conditions}
|
||||||
|
|
||||||
query_vector = self.embedding_model.encode(query).tolist()
|
query_vector = self._encode_cached(query).tolist()
|
||||||
recall_k = RERANK_CANDIDATES if (USE_RERANK or USE_HYBRID_SEARCH) else top_k
|
recall_k = RERANK_CANDIDATES if (USE_RERANK or USE_HYBRID_SEARCH) else top_k
|
||||||
recall_k = max(recall_k, top_k * RECALL_MULTIPLIER)
|
recall_k = max(recall_k, top_k * RECALL_MULTIPLIER)
|
||||||
|
|
||||||
@@ -811,6 +811,76 @@ class RAGEngine:
|
|||||||
logger.warning(f"FAQ 集合查询失败: {e}")
|
logger.warning(f"FAQ 集合查询失败: {e}")
|
||||||
return get_empty_result()
|
return get_empty_result()
|
||||||
|
|
||||||
|
def _encode_cached(self, text):
|
||||||
|
"""
|
||||||
|
带缓存的 embedding 编码
|
||||||
|
|
||||||
|
优先从 Embedding Cache(LRU)读取,未命中再调用模型编码并写入缓存。
|
||||||
|
支持单文本和批量文本输入。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: 单个文本字符串 或 文本列表
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
numpy 数组(单文本为一维,批量为二维)
|
||||||
|
"""
|
||||||
|
import numpy as _np
|
||||||
|
|
||||||
|
# 检查 embedding 缓存是否启用(缓存配置查询结果,避免每次重复导入)
|
||||||
|
if not hasattr(self, '_emb_cache_enabled'):
|
||||||
|
self._emb_cache_enabled = True # 默认启用
|
||||||
|
if CACHE_AVAILABLE:
|
||||||
|
try:
|
||||||
|
from config import EMBEDDING_CACHE_ENABLED
|
||||||
|
self._emb_cache_enabled = EMBEDDING_CACHE_ENABLED
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if not self._emb_cache_enabled:
|
||||||
|
return self.embedding_model.encode(text)
|
||||||
|
|
||||||
|
try:
|
||||||
|
_cache = get_cache_manager()
|
||||||
|
except Exception:
|
||||||
|
return self.embedding_model.encode(text)
|
||||||
|
|
||||||
|
# 批量输入
|
||||||
|
if isinstance(text, list):
|
||||||
|
try:
|
||||||
|
cached_embs, missed_indices = _cache.get_embeddings_batch(text)
|
||||||
|
if missed_indices:
|
||||||
|
missed_texts = [text[i] for i in missed_indices]
|
||||||
|
# encode(list) 始终返回 2D ndarray,直接按行索引即可
|
||||||
|
new_embs = self.embedding_model.encode(missed_texts)
|
||||||
|
if len(missed_indices) == 1:
|
||||||
|
# 单条时 encode 可能返回 1D,需统一处理
|
||||||
|
if new_embs.ndim == 1:
|
||||||
|
new_embs = new_embs.reshape(1, -1)
|
||||||
|
for idx, mi in enumerate(missed_indices):
|
||||||
|
emb_list = new_embs[idx].tolist()
|
||||||
|
cached_embs[mi] = emb_list
|
||||||
|
try:
|
||||||
|
_cache.set_embedding(text[mi], emb_list)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return _np.array(cached_embs)
|
||||||
|
except Exception:
|
||||||
|
# 缓存故障时优雅降级为直接编码
|
||||||
|
return self.embedding_model.encode(text)
|
||||||
|
|
||||||
|
# 单文本输入
|
||||||
|
cached = _cache.get_embedding(text)
|
||||||
|
if cached is not None:
|
||||||
|
return _np.array(cached)
|
||||||
|
|
||||||
|
embedding = self.embedding_model.encode(text)
|
||||||
|
try:
|
||||||
|
emb_list = embedding.tolist() if hasattr(embedding, 'tolist') else list(embedding)
|
||||||
|
_cache.set_embedding(text, emb_list)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return embedding
|
||||||
|
|
||||||
def _search_image_chunks(self, query_vector: list, top_k: int = 5, where_filter: dict = None) -> dict:
|
def _search_image_chunks(self, query_vector: list, top_k: int = 5, where_filter: dict = None) -> dict:
|
||||||
"""
|
"""
|
||||||
独立检索图片切片(P0:图片独立召回通道)
|
独立检索图片切片(P0:图片独立召回通道)
|
||||||
@@ -1171,6 +1241,20 @@ class RAGEngine:
|
|||||||
if not neighbors.get('ids') or len(neighbors.get('ids', [])) <= 1:
|
if not neighbors.get('ids') or len(neighbors.get('ids', [])) <= 1:
|
||||||
neighbors = _get_neighbors({"$and": [{"source": source}, {"chunk_type": "text"}]})
|
neighbors = _get_neighbors({"$and": [{"source": source}, {"chunk_type": "text"}]})
|
||||||
|
|
||||||
|
# 同时扩展同 section 的 table 邻居(table 切片的 rerank 分数往往偏低,
|
||||||
|
# 但与同 section 的 text 切片属于同一语义单元,不应割裂)
|
||||||
|
table_where = {"$and": [{"source": source}, {"chunk_type": "table"}]}
|
||||||
|
if section:
|
||||||
|
table_where["$and"].append({"section": section})
|
||||||
|
table_neighbors = _get_neighbors(table_where)
|
||||||
|
|
||||||
|
# 当 section 为空时,table 查询只有 source 条件,可能拉入大量无关表格,
|
||||||
|
# 缩小 chunk_index 窗口至 ±1 以降低噪音;有 section 时使用正常窗口
|
||||||
|
if section:
|
||||||
|
_t_before, _t_after = CONTEXT_EXPANSION_BEFORE, CONTEXT_EXPANSION_AFTER
|
||||||
|
else:
|
||||||
|
_t_before, _t_after = 1, 1
|
||||||
|
|
||||||
neighbor_rows = []
|
neighbor_rows = []
|
||||||
for n_id, n_doc, n_meta in zip(
|
for n_id, n_doc, n_meta in zip(
|
||||||
neighbors.get('ids', []),
|
neighbors.get('ids', []),
|
||||||
@@ -1183,6 +1267,18 @@ class RAGEngine:
|
|||||||
if seed_index - CONTEXT_EXPANSION_BEFORE <= n_index <= seed_index + CONTEXT_EXPANSION_AFTER:
|
if seed_index - CONTEXT_EXPANSION_BEFORE <= n_index <= seed_index + CONTEXT_EXPANSION_AFTER:
|
||||||
neighbor_rows.append((n_index, n_id, n_doc, n_meta))
|
neighbor_rows.append((n_index, n_id, n_doc, n_meta))
|
||||||
|
|
||||||
|
# 同 section 的 table 邻居也加入扩展范围
|
||||||
|
for n_id, n_doc, n_meta in zip(
|
||||||
|
table_neighbors.get('ids', []),
|
||||||
|
table_neighbors.get('documents', []),
|
||||||
|
table_neighbors.get('metadatas', [])
|
||||||
|
):
|
||||||
|
n_index = self._to_int(n_meta.get('chunk_index'))
|
||||||
|
if n_index is None:
|
||||||
|
continue
|
||||||
|
if seed_index - _t_before <= n_index <= seed_index + _t_after:
|
||||||
|
neighbor_rows.append((n_index, n_id, n_doc, n_meta))
|
||||||
|
|
||||||
seed_neighbors_added = 0
|
seed_neighbors_added = 0
|
||||||
for n_index, n_id, n_doc, n_meta in sorted(neighbor_rows, key=lambda row: row[0]):
|
for n_index, n_id, n_doc, n_meta in sorted(neighbor_rows, key=lambda row: row[0]):
|
||||||
if len(items) >= max_chunks:
|
if len(items) >= max_chunks:
|
||||||
@@ -1387,7 +1483,7 @@ class RAGEngine:
|
|||||||
if not target_collections:
|
if not target_collections:
|
||||||
return get_empty_result()
|
return get_empty_result()
|
||||||
|
|
||||||
query_vector = self.embedding_model.encode(query).tolist()
|
query_vector = self._encode_cached(query).tolist()
|
||||||
# 扩大召回数量,以便过滤废止切片后仍有足够结果
|
# 扩大召回数量,以便过滤废止切片后仍有足够结果
|
||||||
recall_k = RERANK_CANDIDATES if (USE_RERANK or USE_HYBRID_SEARCH) else top_k
|
recall_k = RERANK_CANDIDATES if (USE_RERANK or USE_HYBRID_SEARCH) else top_k
|
||||||
recall_k = max(recall_k, top_k * RECALL_MULTIPLIER)
|
recall_k = max(recall_k, top_k * RECALL_MULTIPLIER)
|
||||||
@@ -1739,12 +1835,12 @@ class RAGEngine:
|
|||||||
# === 高精度版:基于语义向量 ===
|
# === 高精度版:基于语义向量 ===
|
||||||
from core.mmr import mmr_rerank
|
from core.mmr import mmr_rerank
|
||||||
|
|
||||||
# 获取查询向量
|
# 获取查询向量(使用 embedding 缓存)
|
||||||
query_emb = np.array(self.embedding_model.encode(query))
|
query_emb = np.array(self._encode_cached(query))
|
||||||
|
|
||||||
# 批量编码所有文档
|
# 批量编码所有文档(使用 embedding 缓存)
|
||||||
docs_list = results['documents'][0]
|
docs_list = results['documents'][0]
|
||||||
all_embeddings = self.embedding_model.encode(docs_list)
|
all_embeddings = self._encode_cached(docs_list)
|
||||||
|
|
||||||
# 构建候选列表
|
# 构建候选列表
|
||||||
candidates = []
|
candidates = []
|
||||||
@@ -2068,21 +2164,33 @@ class RAGEngine:
|
|||||||
"content": (
|
"content": (
|
||||||
"你是一个严谨的知识库问答助手。"
|
"你是一个严谨的知识库问答助手。"
|
||||||
"你必须且只能根据用户提供的【参考资料】回答问题。"
|
"你必须且只能根据用户提供的【参考资料】回答问题。"
|
||||||
|
"参考资料中每段内容前标有章节路径(━格式),请注意区分不同章节的内容,"
|
||||||
|
"特别当不同章节标题相似或包含相同关键词时,务必根据章节路径准确定位,不要混淆。"
|
||||||
"如果参考资料中有答案,必须引用对应内容回答,并在回答末尾标注引用编号(如[1]、[2])。"
|
"如果参考资料中有答案,必须引用对应内容回答,并在回答末尾标注引用编号(如[1]、[2])。"
|
||||||
"如果参考资料中确实没有相关信息,简短说明即可,不要编造或补充资料外的内容。"
|
"如果参考资料中确实没有相关信息,简短说明即可,不要编造或补充资料外的内容。"
|
||||||
"禁止使用参考资料以外的知识进行补充或推测。"
|
"禁止使用参考资料以外的知识进行补充或推测。"
|
||||||
|
"【重要-表格处理规则】当用户询问表格、要求展示表格内容时,你必须将参考资料中的 Markdown 表格原样输出(保留 | 分隔符和表格结构),"
|
||||||
|
"不要仅用文字描述表格存在或仅列出章节名称。如果参考资料中多个章节都有表格,"
|
||||||
|
"优先展示与用户问题最相关的表格完整内容。"
|
||||||
)
|
)
|
||||||
})
|
})
|
||||||
|
|
||||||
# 添加当前问题(带上下文)- 强化指令
|
# 添加当前问题(带上下文)- 强化指令
|
||||||
if context:
|
if context:
|
||||||
|
# 检测用户问题是否涉及表格,加入针对性指令
|
||||||
|
_table_hint = ""
|
||||||
|
# 检测上下文中是否包含 Markdown 表格(数据驱动,无需硬编码关键词)
|
||||||
|
_has_table_in_context = bool(re.search(r'\|.+\|', context)) if context else False
|
||||||
|
if _has_table_in_context:
|
||||||
|
_table_hint = "\n注意:参考资料中包含 Markdown 格式的表格数据,请务必将相关表格以原始 Markdown 表格格式完整展示在回答中,不要仅用文字描述。"
|
||||||
|
|
||||||
user_message = f"""【参考资料】
|
user_message = f"""【参考资料】
|
||||||
{context}
|
{context}
|
||||||
|
|
||||||
【用户问题】
|
【用户问题】
|
||||||
{query}
|
{query}
|
||||||
|
|
||||||
请仔细阅读以上全部参考资料后回答。如果参考资料中包含相关内容,必须引用回答并标注编号。如果资料中没有相关信息,请明确说明。"""
|
请仔细阅读以上全部参考资料后回答。注意参考资料中标有章节路径,请根据章节路径准确定位相关内容。如果参考资料中包含相关内容,必须引用回答并标注编号。如果资料中没有相关信息,请明确说明。{_table_hint}"""
|
||||||
else:
|
else:
|
||||||
user_message = query
|
user_message = query
|
||||||
|
|
||||||
|
|||||||
@@ -412,6 +412,11 @@ class KnowledgeBaseManager(
|
|||||||
if len(chunks) < 2:
|
if len(chunks) < 2:
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
|
# 检测页码是否可靠:若所有 chunk 的 page_start 相同(如 Word 文档 page_idx 全为 0),
|
||||||
|
# 则页码信息不可用,需要启用降级合并规则
|
||||||
|
page_values = set(getattr(c, 'page_start', 0) for c in chunks)
|
||||||
|
pages_unavailable = len(page_values) <= 1
|
||||||
|
|
||||||
merged_chunks = []
|
merged_chunks = []
|
||||||
i = 0
|
i = 0
|
||||||
merge_count = 0
|
merge_count = 0
|
||||||
@@ -424,6 +429,7 @@ class KnowledgeBaseManager(
|
|||||||
# 查找下一个表格(跳过中间的"续表"文本)
|
# 查找下一个表格(跳过中间的"续表"文本)
|
||||||
next_table_idx = None
|
next_table_idx = None
|
||||||
next_chunk = None
|
next_chunk = None
|
||||||
|
intermediate_texts = [] # 收集中间文本用于降级判断
|
||||||
|
|
||||||
for j in range(i + 1, min(i + 4, len(chunks))): # 最多向前看3个切片
|
for j in range(i + 1, min(i + 4, len(chunks))): # 最多向前看3个切片
|
||||||
candidate = chunks[j]
|
candidate = chunks[j]
|
||||||
@@ -438,7 +444,11 @@ class KnowledgeBaseManager(
|
|||||||
elif candidate_type == 'text' and ('续表' in candidate_title or '续表' in candidate_content):
|
elif candidate_type == 'text' and ('续表' in candidate_title or '续表' in candidate_content):
|
||||||
# 遇到"续表"文本,继续查找下一个表格
|
# 遇到"续表"文本,继续查找下一个表格
|
||||||
continue
|
continue
|
||||||
elif candidate_type not in ('text',):
|
elif candidate_type == 'text':
|
||||||
|
# 非"续表"文本,收集后停止查找
|
||||||
|
intermediate_texts.append(candidate)
|
||||||
|
break
|
||||||
|
else:
|
||||||
# 遇到非文本类型,停止查找
|
# 遇到非文本类型,停止查找
|
||||||
break
|
break
|
||||||
|
|
||||||
@@ -458,24 +468,59 @@ class KnowledgeBaseManager(
|
|||||||
# 获取内容(用于检测"续表")
|
# 获取内容(用于检测"续表")
|
||||||
next_content = getattr(next_chunk, 'content', '')
|
next_content = getattr(next_chunk, 'content', '')
|
||||||
|
|
||||||
|
# 通用/无意义标题集合,这些标题不能用于"标题相似"判定
|
||||||
|
_GENERIC_TITLES = {'表格', 'table', '表格', ''}
|
||||||
|
|
||||||
# 判断是否为跨页表格
|
# 判断是否为跨页表格
|
||||||
is_cross_page = False
|
is_cross_page = False
|
||||||
|
|
||||||
# 规则1: 页码连续(如果页码有效)
|
# 规则1: 页码连续(如果页码有效)
|
||||||
page_valid = curr_page_end > 0 and next_page_start > 0
|
page_valid = curr_page_end > 0 and next_page_start > 0
|
||||||
if page_valid and curr_page_end + 1 == next_page_start:
|
if page_valid and curr_page_end + 1 == next_page_start:
|
||||||
is_cross_page = True
|
# 页码连续时,还需标题匹配或为通用标题才合并
|
||||||
|
# 避免把不同页面上不相关的表格错误合并
|
||||||
|
if curr_title == next_title or curr_title in _GENERIC_TITLES and next_title in _GENERIC_TITLES:
|
||||||
|
is_cross_page = True
|
||||||
|
elif curr_title and next_title:
|
||||||
|
clean_next_r1 = next_title.replace('续表', '').strip()
|
||||||
|
if curr_title in clean_next_r1 or clean_next_r1 in curr_title:
|
||||||
|
is_cross_page = True
|
||||||
|
|
||||||
# 规则2: 第二个表格标题或内容包含"续表"
|
# 规则2: 第二个表格标题或内容包含"续表"
|
||||||
elif '续表' in next_title or '续表' in next_content:
|
elif '续表' in next_title or '续表' in next_content:
|
||||||
is_cross_page = True
|
is_cross_page = True
|
||||||
|
|
||||||
# 规则3: 标题相似(去掉"续表"后比较)
|
# 规则3: 标题相似(去掉"续表"后比较)
|
||||||
elif curr_title and next_title:
|
# 排除通用标题(如"表格"),防止把所有标题为"表格"的相邻表格都误合并
|
||||||
|
elif (curr_title and next_title
|
||||||
|
and curr_title not in _GENERIC_TITLES
|
||||||
|
and next_title not in _GENERIC_TITLES):
|
||||||
clean_next = next_title.replace('续表', '').strip()
|
clean_next = next_title.replace('续表', '').strip()
|
||||||
if curr_title in clean_next or clean_next in curr_title:
|
if clean_next and (curr_title in clean_next or clean_next in curr_title):
|
||||||
is_cross_page = True
|
is_cross_page = True
|
||||||
|
|
||||||
|
# 规则4(降级): 页码不可用(如 Word 文档 page_idx 全为 0)
|
||||||
|
# 仅当页码信息缺失时才启用此规则,避免 PDF 正常页码时被误合并
|
||||||
|
if (not is_cross_page
|
||||||
|
and pages_unavailable
|
||||||
|
and curr_title in _GENERIC_TITLES
|
||||||
|
and next_title in _GENERIC_TITLES):
|
||||||
|
# 检查中间文本是否暗示跨页延续(空、短文本、续表标记等)
|
||||||
|
has_separating_content = False
|
||||||
|
for text_chunk in intermediate_texts:
|
||||||
|
tc = (getattr(text_chunk, 'content', '') or '').strip()
|
||||||
|
tt = (getattr(text_chunk, 'title', '') or '').strip()
|
||||||
|
if not tc:
|
||||||
|
continue # 空文本不算分隔
|
||||||
|
if '续表' in tc or '续表' in tt:
|
||||||
|
continue # 续表标记,说明是跨页
|
||||||
|
# 有实质性中间内容(如分类标题"A3类:xxx"),不合并
|
||||||
|
has_separating_content = True
|
||||||
|
break
|
||||||
|
if not has_separating_content:
|
||||||
|
is_cross_page = True
|
||||||
|
logger.debug(f"降级合并(页码不可用): '{curr_title}' + '{next_title}'")
|
||||||
|
|
||||||
if is_cross_page:
|
if is_cross_page:
|
||||||
# 执行合并
|
# 执行合并
|
||||||
merge_count += 1
|
merge_count += 1
|
||||||
@@ -488,14 +533,27 @@ class KnowledgeBaseManager(
|
|||||||
# 合并两个表格的 HTML
|
# 合并两个表格的 HTML
|
||||||
current.table_html = curr_html + '\n' + next_html
|
current.table_html = curr_html + '\n' + next_html
|
||||||
|
|
||||||
# 合并 image_path 到 images
|
# 合并 image_path 和嵌入图片到 images
|
||||||
curr_img = getattr(current, 'image_path', None)
|
curr_img = getattr(current, 'image_path', None)
|
||||||
next_img = getattr(next_chunk, 'image_path', None)
|
next_img = getattr(next_chunk, 'image_path', None)
|
||||||
merged_images = []
|
curr_images = getattr(current, 'images', None) or []
|
||||||
if curr_img:
|
next_images = getattr(next_chunk, 'images', None) or []
|
||||||
|
|
||||||
|
# 合并两个表格的所有图片(image_path + 嵌入图片)
|
||||||
|
merged_images = list(curr_images) # 保留当前表格的嵌入图片
|
||||||
|
# 添加 image_path 图片(如果不在列表中)
|
||||||
|
existing_ids = {img.get('id', '') for img in merged_images if isinstance(img, dict)}
|
||||||
|
if curr_img and curr_img not in existing_ids:
|
||||||
merged_images.append({'id': curr_img, 'page': curr_page_end})
|
merged_images.append({'id': curr_img, 'page': curr_page_end})
|
||||||
if next_img:
|
existing_ids.add(curr_img)
|
||||||
|
for img in next_images: # 添加下一个表格的嵌入图片
|
||||||
|
img_id = img.get('id', '') if isinstance(img, dict) else ''
|
||||||
|
if img_id and img_id not in existing_ids:
|
||||||
|
merged_images.append(img)
|
||||||
|
existing_ids.add(img_id)
|
||||||
|
if next_img and next_img not in existing_ids:
|
||||||
merged_images.append({'id': next_img, 'page': next_page_start})
|
merged_images.append({'id': next_img, 'page': next_page_start})
|
||||||
|
|
||||||
if merged_images:
|
if merged_images:
|
||||||
current.images = merged_images
|
current.images = merged_images
|
||||||
# 保留第一个图片作为主 image_path
|
# 保留第一个图片作为主 image_path
|
||||||
@@ -544,7 +602,7 @@ class KnowledgeBaseManager(
|
|||||||
try:
|
try:
|
||||||
from config import get_llm_client, DASHSCOPE_MODEL
|
from config import get_llm_client, DASHSCOPE_MODEL
|
||||||
client = get_llm_client()
|
client = get_llm_client()
|
||||||
summary = call_llm(client, prompt, DASHSCOPE_MODEL, max_tokens=100)
|
summary = call_llm(client, prompt, DASHSCOPE_MODEL, max_tokens=512)
|
||||||
return summary.strip() if summary else ""
|
return summary.strip() if summary else ""
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"生成表格摘要失败: {e}")
|
logger.warning(f"生成表格摘要失败: {e}")
|
||||||
@@ -603,7 +661,7 @@ class KnowledgeBaseManager(
|
|||||||
]
|
]
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
max_tokens=200
|
max_tokens=512
|
||||||
)
|
)
|
||||||
|
|
||||||
description = response.choices[0].message.content
|
description = response.choices[0].message.content
|
||||||
|
|||||||
@@ -283,7 +283,7 @@ class KnowledgeBaseRouter:
|
|||||||
content = call_llm(
|
content = call_llm(
|
||||||
self.llm_client, prompt, MODEL,
|
self.llm_client, prompt, MODEL,
|
||||||
temperature=0.1,
|
temperature=0.1,
|
||||||
max_tokens=100
|
max_tokens=512
|
||||||
)
|
)
|
||||||
|
|
||||||
if content is None:
|
if content is None:
|
||||||
|
|||||||
347
parsers/heading_rules.py
Normal file
347
parsers/heading_rules.py
Normal file
@@ -0,0 +1,347 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""
|
||||||
|
标题识别规则引擎
|
||||||
|
|
||||||
|
将 _detect_heading_level 的硬编码正则提取为可配置的规则列表。
|
||||||
|
规则按优先级从高到低排序,第一个匹配即返回。
|
||||||
|
|
||||||
|
MinerU 解析 DOCX 等 Office 格式时通常不提供 text_level(全部为 0),
|
||||||
|
此时需要启发式识别标题层级。本模块提供可配置的规则引擎替代原来的
|
||||||
|
硬编码 if-elif 链。
|
||||||
|
|
||||||
|
设计要点:
|
||||||
|
- HeadingRule 数据类支持正向匹配(pattern)和反向排除(exclude_pattern)
|
||||||
|
- 长度约束(min_length / max_length)可精确控制匹配范围
|
||||||
|
- 规则可单独禁用(enabled=False),便于调试
|
||||||
|
- 全局单例通过 config.py 覆盖默认值
|
||||||
|
|
||||||
|
MinerU v2 格式备注:
|
||||||
|
content_list_v2.json 中的 paragraph_content 包含 style=["bold"] 信息,
|
||||||
|
layout.json 中的 spans 也有 style 信息。这些信息比正则匹配 **加粗** 更可靠,
|
||||||
|
但当前代码使用 v1 格式(content_list.json),暂不利用 v2 的 style。
|
||||||
|
HeadingRuleEngine.detect() 签名预留了 style 参数,未来切换到 v2 格式后
|
||||||
|
可直接利用 style 信息辅助判断。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import re
|
||||||
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional, List, Tuple
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class HeadingRule:
|
||||||
|
"""
|
||||||
|
标题识别规则
|
||||||
|
|
||||||
|
每条规则定义一个文本模式到标题级别的映射。
|
||||||
|
规则引擎按列表顺序逐条匹配,第一个命中即返回。
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
pattern: 编译后的正则(match 语义,从文本开头匹配)
|
||||||
|
level: 匹配时返回的标题级别 (1=h1, 2=h2, 3=h3)
|
||||||
|
name: 规则名称(用于日志和配置覆盖)
|
||||||
|
max_length: 文本最大长度,0=不限
|
||||||
|
min_length: 文本最小长度,0=不限
|
||||||
|
enabled: 是否启用
|
||||||
|
exclude_pattern: 匹配此模式则排除(反向过滤)
|
||||||
|
|
||||||
|
Example:
|
||||||
|
>>> rule = HeadingRule(
|
||||||
|
... pattern=re.compile(r'^第[一二三四五六七八九十百千万]+[章节篇部]'),
|
||||||
|
... level=1,
|
||||||
|
... name="chinese_chapter",
|
||||||
|
... )
|
||||||
|
>>> rule.match("第一章 总则")
|
||||||
|
1
|
||||||
|
>>> rule.match("这是正文")
|
||||||
|
0
|
||||||
|
"""
|
||||||
|
|
||||||
|
pattern: re.Pattern
|
||||||
|
level: int
|
||||||
|
name: str
|
||||||
|
max_length: int = 0
|
||||||
|
min_length: int = 0
|
||||||
|
enabled: bool = True
|
||||||
|
exclude_pattern: Optional[re.Pattern] = None
|
||||||
|
|
||||||
|
def match(self, text: str) -> int:
|
||||||
|
"""
|
||||||
|
检查文本是否匹配此规则
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: 待检测文本(调用前应已 strip)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
标题级别,0 表示不匹配
|
||||||
|
"""
|
||||||
|
if not self.enabled:
|
||||||
|
return 0
|
||||||
|
if self.min_length > 0 and len(text) < self.min_length:
|
||||||
|
return 0
|
||||||
|
if self.max_length > 0 and len(text) > self.max_length:
|
||||||
|
return 0
|
||||||
|
if self.exclude_pattern and self.exclude_pattern.search(text):
|
||||||
|
return 0
|
||||||
|
if self.pattern.match(text):
|
||||||
|
return self.level
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
# 默认规则列表(按优先级从高到低)
|
||||||
|
#
|
||||||
|
# 注意事项:
|
||||||
|
# - 数字三级标题 (1.1.1) 必须在二级 (1.1) 之前,因为 1.1.1 也匹配 ^\d+\.\d+
|
||||||
|
# - short_chinese_heading 是最宽泛的规则,放在最后作为兜底
|
||||||
|
# - 第 9 条规则相比原版增加了 exclude_pattern,排除以句末标点结尾的短文本
|
||||||
|
DEFAULT_HEADING_RULES: List[HeadingRule] = [
|
||||||
|
# 1. 中文章节标题 -> h1
|
||||||
|
# 匹配:第一章、第二章、第十节、第三篇 等
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^第[一二三四五六七八九十百千万]+[章节篇部]'),
|
||||||
|
level=1,
|
||||||
|
name="chinese_chapter",
|
||||||
|
),
|
||||||
|
# 2. 中文条款编号 -> h2
|
||||||
|
# 匹配:第一条、第三款 等
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^第[一二三四五六七八九十百千万]+[条款]'),
|
||||||
|
level=2,
|
||||||
|
name="chinese_article",
|
||||||
|
),
|
||||||
|
# 3. 数字三级标题 -> h3(必须在二级之前匹配)
|
||||||
|
# 匹配:1.1.1 背景、2.3.4 方案 等
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^\d+\.\d+\.\d+[\.、\s]'),
|
||||||
|
level=3,
|
||||||
|
name="numeric_level3",
|
||||||
|
max_length=100,
|
||||||
|
),
|
||||||
|
# 4. 数字二级标题 -> h2(必须在一级之前匹配)
|
||||||
|
# 匹配:1.1 背景、2.3 方案 等
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^\d+\.\d+[\.、\s]'),
|
||||||
|
level=2,
|
||||||
|
name="numeric_level2",
|
||||||
|
max_length=80,
|
||||||
|
),
|
||||||
|
# 5. 数字一级标题 -> h1
|
||||||
|
# 匹配:1. 概述、2、背景 等
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^\d+[\.、\s]'),
|
||||||
|
level=1,
|
||||||
|
name="numeric_level1",
|
||||||
|
max_length=50,
|
||||||
|
),
|
||||||
|
# 6. 英文章节标题 -> h1
|
||||||
|
# 匹配:Chapter 1、Section 2、Part 3 等
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^(Chapter|Section|Part|Chapter\s+\d+|Section\s+\d+)', re.IGNORECASE),
|
||||||
|
level=1,
|
||||||
|
name="english_chapter",
|
||||||
|
),
|
||||||
|
# 7. 分类标题 -> h3(必须在 bold_short_text 之前,否则 **A2类:** 会被加粗规则抢先匹配)
|
||||||
|
# 匹配:A1类:公园、**A2类**:各类卫生医疗机构、**B1类:** 道路 等
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^\*{0,2}[A-Z]\d+[类類]\*{0,2}[::]'),
|
||||||
|
level=3,
|
||||||
|
name="category_heading",
|
||||||
|
),
|
||||||
|
# 8. 加粗短文本 -> h2
|
||||||
|
# 匹配:**重要通知**、**概述** 等(Markdown 加粗标记)
|
||||||
|
# 注意:**A2类:** 已被分类标题规则优先匹配,不会误判为 h2
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'^\*\*.+\*\*$'),
|
||||||
|
level=2,
|
||||||
|
name="bold_short_text",
|
||||||
|
max_length=50,
|
||||||
|
),
|
||||||
|
# 9. 短中文文本 -> h2(替代原"任何 <20 字符含中文"规则)
|
||||||
|
# 关键改进:排除以句末标点结尾的文本
|
||||||
|
# 原规则将 "这是一段正文。" 也识别为 h2,导致大量误判
|
||||||
|
# 新规则:包含中文 + 长度 2-20 + 不以句末标点结尾 → h2
|
||||||
|
HeadingRule(
|
||||||
|
pattern=re.compile(r'[一-鿿]'),
|
||||||
|
level=2,
|
||||||
|
name="short_chinese_heading",
|
||||||
|
max_length=20,
|
||||||
|
min_length=2,
|
||||||
|
exclude_pattern=re.compile(r'[。!?;…]$'),
|
||||||
|
enabled=True,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class HeadingRuleEngine:
|
||||||
|
"""
|
||||||
|
标题识别规则引擎
|
||||||
|
|
||||||
|
按规则列表顺序逐条匹配,第一个命中即返回标题级别。
|
||||||
|
支持从 config.py 加载自定义规则或覆盖默认规则参数。
|
||||||
|
|
||||||
|
Example:
|
||||||
|
>>> engine = HeadingRuleEngine()
|
||||||
|
>>> engine.detect("第一章 总则")
|
||||||
|
(1, 'chinese_chapter')
|
||||||
|
>>> engine.detect("这是普通正文。")
|
||||||
|
(0, None)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, rules: Optional[List[HeadingRule]] = None) -> None:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
rules: 规则列表,None 则使用默认规则的深拷贝
|
||||||
|
"""
|
||||||
|
if rules is not None:
|
||||||
|
self.rules: List[HeadingRule] = rules
|
||||||
|
else:
|
||||||
|
import copy
|
||||||
|
self.rules = copy.deepcopy(DEFAULT_HEADING_RULES)
|
||||||
|
|
||||||
|
def detect(self, text: str, style: Optional[List[str]] = None) -> Tuple[int, Optional[str]]:
|
||||||
|
"""
|
||||||
|
检测文本的标题级别
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: 待检测文本
|
||||||
|
style: MinerU v2 格式中的 style 信息(如 ["bold"]),
|
||||||
|
当文本标记为 bold 且较短时,可直接判定为标题,
|
||||||
|
无需依赖 Markdown **...** 标记。
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(level, rule_name): 标题级别和匹配的规则名
|
||||||
|
level=0 表示不是标题
|
||||||
|
"""
|
||||||
|
text = text.strip()
|
||||||
|
if not text:
|
||||||
|
return 0, None
|
||||||
|
|
||||||
|
# v2 style 信息:如果文本标记为 bold 且较短,优先尝试加粗规则
|
||||||
|
if style and 'bold' in style and 2 <= len(text) <= 50:
|
||||||
|
# 先检查是否匹配更高优先级的分类标题规则
|
||||||
|
for rule in self.rules:
|
||||||
|
if rule.name == 'category_heading' and rule.enabled:
|
||||||
|
level = rule.match(text)
|
||||||
|
if level > 0:
|
||||||
|
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||||
|
return level, rule.name
|
||||||
|
|
||||||
|
# 再检查是否匹配中文章节/条款等高优先级规则
|
||||||
|
for rule in self.rules:
|
||||||
|
if rule.name in ('chinese_chapter', 'chinese_article', 'numeric_level3',
|
||||||
|
'numeric_level2', 'numeric_level1', 'english_chapter') and rule.enabled:
|
||||||
|
level = rule.match(text)
|
||||||
|
if level > 0:
|
||||||
|
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||||
|
return level, rule.name
|
||||||
|
|
||||||
|
# 否则作为加粗短文本 → h2(与 bold_short_text 规则对齐,但不依赖 **...** 标记)
|
||||||
|
logger.debug(f"标题识别(v2 style): '{text[:30]}' -> h2 (规则: bold_short_text_via_style)")
|
||||||
|
return 2, 'bold_short_text'
|
||||||
|
|
||||||
|
# 常规规则匹配(v1 格式或无 style 信息时)
|
||||||
|
for rule in self.rules:
|
||||||
|
level = rule.match(text)
|
||||||
|
if level > 0:
|
||||||
|
logger.debug(f"标题识别: '{text[:30]}' -> h{level} (规则: {rule.name})")
|
||||||
|
return level, rule.name
|
||||||
|
|
||||||
|
return 0, None
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== 全局单例 ====================
|
||||||
|
|
||||||
|
_engine: Optional[HeadingRuleEngine] = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_heading_engine() -> HeadingRuleEngine:
|
||||||
|
"""获取全局标题识别引擎(延迟初始化,线程安全)"""
|
||||||
|
global _engine
|
||||||
|
if _engine is None:
|
||||||
|
_engine = _create_engine_from_config()
|
||||||
|
return _engine
|
||||||
|
|
||||||
|
|
||||||
|
def _create_engine_from_config() -> HeadingRuleEngine:
|
||||||
|
"""
|
||||||
|
从 config 创建引擎(支持配置覆盖)
|
||||||
|
|
||||||
|
优先级:
|
||||||
|
1. config.HEADING_RULES_CONFIG 不为 None → 使用自定义规则
|
||||||
|
2. config 细粒度参数覆盖默认规则(如 HEADING_SHORT_TEXT_ENABLED)
|
||||||
|
3. 使用默认规则
|
||||||
|
"""
|
||||||
|
# 尝试加载完整自定义规则
|
||||||
|
try:
|
||||||
|
from config import HEADING_RULES_CONFIG
|
||||||
|
if HEADING_RULES_CONFIG is not None:
|
||||||
|
rules = _build_rules_from_config(HEADING_RULES_CONFIG)
|
||||||
|
logger.info(f"使用自定义标题规则: {len(rules)} 条")
|
||||||
|
return HeadingRuleEngine(rules)
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# 使用默认规则,应用细粒度配置覆盖
|
||||||
|
rules = list(DEFAULT_HEADING_RULES)
|
||||||
|
try:
|
||||||
|
from config import HEADING_SHORT_TEXT_ENABLED
|
||||||
|
for rule in rules:
|
||||||
|
if rule.name == "short_chinese_heading":
|
||||||
|
rule.enabled = HEADING_SHORT_TEXT_ENABLED
|
||||||
|
logger.debug(f"配置覆盖: short_chinese_heading.enabled={HEADING_SHORT_TEXT_ENABLED}")
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
try:
|
||||||
|
from config import HEADING_SHORT_TEXT_MAX_LENGTH
|
||||||
|
for rule in rules:
|
||||||
|
if rule.name == "short_chinese_heading":
|
||||||
|
rule.max_length = HEADING_SHORT_TEXT_MAX_LENGTH
|
||||||
|
logger.debug(f"配置覆盖: short_chinese_heading.max_length={HEADING_SHORT_TEXT_MAX_LENGTH}")
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return HeadingRuleEngine(rules)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_rules_from_config(config: list) -> List[HeadingRule]:
|
||||||
|
"""
|
||||||
|
从配置字典列表构建规则列表
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: 规则配置列表,每项为 dict,包含:
|
||||||
|
- pattern (str): 正则表达式字符串
|
||||||
|
- level (int): 标题级别
|
||||||
|
- name (str): 规则名称
|
||||||
|
- max_length (int, 可选): 文本最大长度
|
||||||
|
- min_length (int, 可选): 文本最小长度
|
||||||
|
- enabled (bool, 可选): 是否启用
|
||||||
|
- exclude_pattern (str, 可选): 排除正则
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
规则列表
|
||||||
|
"""
|
||||||
|
rules = []
|
||||||
|
for item in config:
|
||||||
|
exclude = None
|
||||||
|
if 'exclude_pattern' in item:
|
||||||
|
exclude = re.compile(item['exclude_pattern'])
|
||||||
|
rules.append(HeadingRule(
|
||||||
|
pattern=re.compile(item['pattern']),
|
||||||
|
level=item['level'],
|
||||||
|
name=item['name'],
|
||||||
|
max_length=item.get('max_length', 0),
|
||||||
|
min_length=item.get('min_length', 0),
|
||||||
|
enabled=item.get('enabled', True),
|
||||||
|
exclude_pattern=exclude,
|
||||||
|
))
|
||||||
|
return rules
|
||||||
|
|
||||||
|
|
||||||
|
def reset_heading_engine() -> None:
|
||||||
|
"""重置引擎(用于测试)"""
|
||||||
|
global _engine
|
||||||
|
_engine = None
|
||||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user