优化切片,增加章节图片信息

This commit is contained in:
Defeng 2026-07-20 20:58:54 +08:00
parent 3ffa6704ad
commit 5a2ee98dfb

View File

@ -342,6 +342,15 @@ def record_to_chunk_text(ins: Dict[str, Any]) -> str:
ins.get("image_footnote"), ins.get("image_footnote"),
) )
if record_type == "chart":
return join_nonempty_parts(
ocr_text,
ins.get("chart_caption"),
ins.get("content"),
ins.get("img_path"),
ins.get("chart_footnote"),
)
if record_type == "equation": if record_type == "equation":
return join_nonempty_parts( return join_nonempty_parts(
ocr_text, ocr_text,
@ -371,6 +380,7 @@ def is_supported_chunk_record(ins: Dict[str, Any]) -> bool:
"text", "text",
"table", "table",
"image", "image",
"chart",
"equation", "equation",
"code", "code",
"list", "list",
@ -385,6 +395,7 @@ def append_caption_ocr(ins: Dict[str, Any], ocr_texts: List[str]) -> None:
caption_field = { caption_field = {
"image": "image_caption", "image": "image_caption",
"table": "table_caption", "table": "table_caption",
"chart": "chart_caption",
"code": "code_caption", "code": "code_caption",
}.get(ins.get("type")) }.get(ins.get("type"))
if not caption_field: if not caption_field:
@ -570,10 +581,10 @@ def extract_maintenance_groups_and_remaining(
) )
if is_start: if is_start:
# 如果是第一个维修表,跳过(不加入 groups也不标记 used_ids # 如果是第一个维修表,跳过(不加入 groups也不标记 used_ids
if not skipped_first: if not skipped_first:
skipped_first = True skipped_first = True
# 注意:这里不清空 current_group因为前面不可能有 group第一个 start # 注意:这里不清空 current_group因为前面不可能有 group第一个 start
current_group = None # 确保不会把之前的非 start 内容误加 current_group = None # 确保不会把之前的非 start 内容误加
continue # 跳过这个 record不加入任何 group continue # 跳过这个 record不加入任何 group
@ -586,7 +597,7 @@ def extract_maintenance_groups_and_remaining(
if current_group is not None: if current_group is not None:
current_group.append(record) current_group.append(record)
# 处理最后一个 group仅当有有效 group 时) # 处理最后一个 group仅当有有效 group 时)
if current_group is not None: if current_group is not None:
groups.append(current_group) groups.append(current_group)
used_ids.update(r["id"] for r in current_group) used_ids.update(r["id"] for r in current_group)
@ -635,7 +646,7 @@ def is_operation_table_by_first_tr(table_body: str) -> bool:
def extract_operation_groups_and_remaining(records, max_following_for_last=23): def extract_operation_groups_and_remaining(records, max_following_for_last=23):
""" """
通过语义判断表格是否为操作项目表基于第一行内容并分组 通过语义判断表格是否为操作项目表基于第一行内容并分组
注意第一个识别到的操作项目表及其后续记录组成的 group 会被跳过不加入 groups 注意第一个识别到的操作项目表及其后续记录组成的 group 会被跳过不加入 groups
但这些记录仍保留在 remaining_records 但这些记录仍保留在 remaining_records
普通起始标记 = [start, ..., next_start - 1] 普通起始标记 = [start, ..., next_start - 1]
@ -656,7 +667,7 @@ def extract_operation_groups_and_remaining(records, max_following_for_last=23):
if current_group is not None: if current_group is not None:
# 决定是否保留当前积累的 group # 决定是否保留当前积累的 group
if not first_group_skipped: if not first_group_skipped:
# 跳过第一个 group不清空 current_group也不加入 groups 或 used_ids # 跳过第一个 group不清空 current_group也不加入 groups 或 used_ids
first_group_skipped = True first_group_skipped = True
# 注意:这里不把 current_group 加入 groups也不更新 used_ids # 注意:这里不把 current_group 加入 groups也不更新 used_ids
else: else:
@ -671,7 +682,7 @@ def extract_operation_groups_and_remaining(records, max_following_for_last=23):
if current_group is not None: if current_group is not None:
current_group.append(record) current_group.append(record)
# 处理最后一个 group如果存在 # 处理最后一个 group如果存在
if current_group is not None: if current_group is not None:
if not first_group_skipped: if not first_group_skipped:
# 整个数据中只有一个 group且是第一个 → 跳过它 # 整个数据中只有一个 group且是第一个 → 跳过它
@ -680,7 +691,7 @@ def extract_operation_groups_and_remaining(records, max_following_for_last=23):
groups.append(current_group) groups.append(current_group)
used_ids.update(r["id"] for r in current_group) used_ids.update(r["id"] for r in current_group)
# —————— 后处理:截断最后一个 group如果需要 —————— # —————— 后处理:截断最后一个 group如果需要 ——————
if groups and max_following_for_last is not None: if groups and max_following_for_last is not None:
last_group = groups[-1] last_group = groups[-1]
if len(last_group) > max_following_for_last + 1: # +1 是起始记录本身 if len(last_group) > max_following_for_last + 1: # +1 是起始记录本身
@ -729,7 +740,7 @@ def extract_wxanli(records):
# 收集当前切片 # 收集当前切片
slice_records = [] slice_records = []
# 添加前2个记录如果存在且未超出边界 # 添加前2个记录如果存在且未超出边界
for offset in [2, 1]: for offset in [2, 1]:
prev_idx = start_pos - offset prev_idx = start_pos - offset
if prev_idx >= 0: if prev_idx >= 0:
@ -810,85 +821,67 @@ def split_content_bbox(data, max_length=8000):
# def group_records(records): def _group_records_without_parent_context(records):
# result = []
# current_group = []
# for i, ins in enumerate(records):
# if 'text_level' in ins and ins['text_level']:
# # 当遇到新的text_level时如果current_group非空则先将其添加到结果中
# if current_group:
# result.append(current_group)
# current_group = [] # 开始新一组
# current_group.append(ins)
# else:
# # 如果当前记录没有text_level或其值为空则继续添加到当前组
# current_group.append(ins)
# # 别忘了将最后一个组添加到结果中
# if current_group:
# result.append(current_group)
# return result
def group_records(records):
"""
text_level 分组并为每个组合适的父级标题上下文
规则
- text_level=N (N>1) 的组 在组开头注入最近的 text_level=1 N-1 的标题
"""
result = [] result = []
current_group = [] current_group = []
# 维护标题栈level -> {text, page_idx, bbox}
heading_stack = {}
for i, ins in enumerate(records): for i, ins in enumerate(records):
text_level = ins.get('text_level') if 'text_level' in ins and ins['text_level']:
# 当遇到新的text_level时如果current_group非空则先将其添加到结果中
if text_level is not None and isinstance(text_level, int) and text_level > 0:
# 更新标题栈
heading_stack[text_level] = {
'text': ins.get('text', ''),
'page_idx': ins.get('page_idx', -1),
'bbox': ins.get('bbox', []),
'level': text_level
}
# 清除更低层级的标题(高层级变化后,低层级失效)
levels_to_remove = [l for l in list(heading_stack.keys()) if l > text_level]
for l in levels_to_remove:
del heading_stack[l]
# 遇到新标题,先保存当前组
if current_group: if current_group:
result.append(current_group) result.append(current_group)
current_group = [] current_group = [] # 开始新一组
current_group.append(ins)
# 注入父级标题(如果当前层级 > 1
if text_level > 1:
prefix_records = []
for level in sorted(heading_stack.keys()):
if level < text_level:
h = heading_stack[level]
prefix_records.append({
'type': 'text',
'text': h['text'],
'text_level': level,
'page_idx': h['page_idx'],
'bbox': h['bbox']
})
current_group = prefix_records + [ins]
else:
current_group = [ins]
else: else:
# 非标题记录,继续添加到当前组 # 如果当前记录没有text_level或其值为空则继续添加到当前组
current_group.append(ins) current_group.append(ins)
# 别忘了将最后一个组添加到结果中
if current_group: if current_group:
result.append(current_group) result.append(current_group)
return result return result
def _get_group_level(group):
if not group:
return None
level = group[0].get("text_level")
if isinstance(level, int) and level > 0:
return level
return None
def group_records(records):
"""
Split records by text_level, then prepend ancestor section content.
For example, a text_level=2 section will include the nearest text_level=1
group content before its own content. A text_level=3 section will include
the nearest level 1 and level 2 group content before its own content.
"""
base_groups = _group_records_without_parent_context(records)
result = []
section_stack = [] # [(level, merged_group)]
for group in base_groups:
level = _get_group_level(group)
if level is None:
result.append(group)
continue
while section_stack and section_stack[-1][0] >= level:
section_stack.pop()
parent_group = section_stack[-1][1] if section_stack else []
merged_group = parent_group + group
result.append(merged_group)
section_stack.append((level, merged_group))
return result
def get_chunk_bbox_otherfile(data, max_length=8000): def get_chunk_bbox_otherfile(data, max_length=8000):
""" """
处理没有 bbox 信息的文件数据 处理没有 bbox 信息的文件数据
@ -1120,7 +1113,7 @@ async def data_replace(data, prefix):
def filter_content_bbox(data, max_length=3000): def filter_content_bbox(data, max_length=3000):
""" """
data 中的 text/table/image 内容按字符长度分块每块不超过 max_length data 中的 text/table/image 内容按字符长度分块每块不超过 max_length
返回 [{"content": str, "highlight_list": list}, ...] 返回 [{"content": str, "highlight_list": list}, ...]
""" """
chunks = [] chunks = []
@ -1159,7 +1152,7 @@ def filter_content_bbox(data, max_length=3000):
current_content = "" current_content = ""
current_highlight_list = [] current_highlight_list = []
# 如果单个 added_content 超过 max_length仍需保留避免丢弃 # 如果单个 added_content 超过 max_length仍需保留避免丢弃
if len(added_content) > max_length: if len(added_content) > max_length:
# 强制截断或直接放入(这里选择直接放入,避免信息丢失) # 强制截断或直接放入(这里选择直接放入,避免信息丢失)
current_content = added_content[:max_length] current_content = added_content[:max_length]
@ -1260,9 +1253,9 @@ def find_ship_info_by_hull(json_file_path, data):
print(f"错误:找不到文件 {json_file_path}") print(f"错误:找不到文件 {json_file_path}")
return None return None
except json.JSONDecodeError: except json.JSONDecodeError:
print("错误JSON 文件格式不正确") print("错误JSON 文件格式不正确")
return None return None
# def merge_short_slices(slices, min_length=30,filename="122-06A0014-B01003_雷达-使用说明书.pdf"): def merge_short_slices(slices, min_length=30,filename="122-06A0014-B01003_雷达-使用说明书.pdf"):
""" """
合并过短的切片 合并过短的切片
- 如果某切片 content 长度 <= min_length则将其合并到下一个切片的开头 - 如果某切片 content 长度 <= min_length则将其合并到下一个切片的开头
@ -1358,110 +1351,6 @@ def find_ship_info_by_hull(json_file_path, data):
ins['content'] = content[:newline_idx] + suffix + content[newline_idx:] ins['content'] = content[:newline_idx] + suffix + content[newline_idx:]
return result return result
def merge_short_slices(slices, min_length=30, filename="122-06A0014-B01003_雷达-使用说明书.pdf"):
"""
合并过短的切片
- 如果某切片 content 长度 <= min_length则将其合并到下一个切片的开头
- 若处于末尾无下一个切片则反向合并到上一个切片末尾
- content 用换行拼接positions 顺序拼接
Args:
slices: [{"content": str, "positions": [...]}]
min_length: 短切片的字符长度阈值
filename: 用于实体提取的文件名
Returns:
合并后的 slices 列表
"""
if not slices:
return slices
result = []
pending_contents = [] # 缓存等待合并到"下一个"的短切片 content
pending_positions = [] # 缓存对应的 positions
for ins in slices:
content = ins.get("content", "") or ""
positions = ins.get("positions", []) or []
if len(content) <= min_length:
# 暂存,等到下一个正常长度的切片再合并
pending_contents.append(content)
pending_positions.extend(positions)
else:
# 正常切片:把暂存的短切片合并到它的开头
if pending_contents:
merged_prefix = "\n".join(pending_contents)
content = merged_prefix + ("\n" if merged_prefix else "") + content
positions = pending_positions + positions
pending_contents = []
pending_positions = []
result.append({
"content": content,
"positions": positions
})
# 收尾:如果末尾还有未合并的短切片(后面没有正常切片可合并)
# 则反向合并到上一个切片末尾
if pending_contents:
merged_suffix = "\n".join(pending_contents)
if result:
last = result[-1]
last["content"] = (last["content"] or "") + ("\n" if last["content"] else "") + merged_suffix
last["positions"] = (last["positions"] or []) + pending_positions
else:
# 极端情况:所有切片都很短,整体作为一个切片返回
result.append({
"content": merged_suffix,
"positions": pending_positions
})
try:
final_result = get_entity(filename)
xinghao = find_ship_info_by_hull(SHIP_MODEL_NAME, final_result)
if xinghao is None: # 确保 xinghao 为 None 时不会报错
xinghao = {"model_name": "", "ship_name": ""}
print("未找到匹配的舰船信息")
# 安全处理 xinghao 为 None 的情况
model_name = ""
if xinghao and 'model_name' in xinghao:
model_name = xinghao['model_name']
other_data = format_entity_text(final_result)
logger.info(f"文件名实体提取成功: {other_data}")
except Exception as e:
logger.warning(f"文件名实体提取失败: {e}")
other_data = ""
model_name = "" # 确保异常时 model_name 有定义
# 将提取的信息(如舰艇名、型号)注入到每个切片的开头
for ins in result:
content = ins['content']
# 如果没有提取到有效数据,则跳过注入
newline_idx = content.find('\n')
info = other_data
if model_name:
info += f',型号为{model_name}'
# ===== 修改部分:拼接文件名 + 实体信息 =====
# 构建前缀:#文件名为xxx.pdf(实体信息)
filename_prefix = f"#文件名为:{filename}"
if info:
filename_prefix += f"({info}"
# 注意:这里使用全角右括号""以匹配你的示例,如需半角请改为 ")"
# 将前缀插入到 content 的最开头
if content:
ins['content'] = filename_prefix + "\n" + content
else:
ins['content'] = filename_prefix
# ==========================================
return result
def find_and_read_content_list(directory, original_filename, encoding='utf-8'): def find_and_read_content_list(directory, original_filename, encoding='utf-8'):
""" """
根据原始文件名在指定目录及其子目录下查找对应的 _content_list.json 文件并返回内容 根据原始文件名在指定目录及其子目录下查找对应的 _content_list.json 文件并返回内容
@ -1527,6 +1416,8 @@ async def process_pdf_file(pdf_path: Path, image_prefix: str,filename:str) -> Op
slices = get_chunk_bbox(replaced_content_list) slices = get_chunk_bbox(replaced_content_list)
slices_check = chunk_check(slices, 8000) slices_check = chunk_check(slices, 8000)
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename) slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
print(11111111111111111111111111111111111111)
print(len(images))
return {"slices": slices_check, "images": images} return {"slices": slices_check, "images": images}
except json.JSONDecodeError: except json.JSONDecodeError:
@ -1596,7 +1487,7 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) -
if __name__ == "__main__": if __name__ == "__main__":
filepath ="/app/mineru_output/舰船抗沉损管训练仿真系统研究_content_list.json" filepath = r"E:\ZKYNLP\Hjunproject\project0506\kgrag\舰船抗沉损管训练仿真系统研究_content_list.json"
try: try:
with open(filepath, 'r', encoding='utf-8') as file: with open(filepath, 'r', encoding='utf-8') as file:
data = json.load(file) data = json.load(file)
@ -1614,4 +1505,4 @@ if __name__ == "__main__":
for ins in slices_2[:30]: for ins in slices_2[:30]:
print(ins) print(ins)