diff --git a/chunk_text.py b/chunk_text.py index a237670..c47b18c 100644 --- a/chunk_text.py +++ b/chunk_text.py @@ -342,6 +342,15 @@ def record_to_chunk_text(ins: Dict[str, Any]) -> str: 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": return join_nonempty_parts( ocr_text, @@ -371,6 +380,7 @@ def is_supported_chunk_record(ins: Dict[str, Any]) -> bool: "text", "table", "image", + "chart", "equation", "code", "list", @@ -385,6 +395,7 @@ def append_caption_ocr(ins: Dict[str, Any], ocr_texts: List[str]) -> None: caption_field = { "image": "image_caption", "table": "table_caption", + "chart": "chart_caption", "code": "code_caption", }.get(ins.get("type")) if not caption_field: @@ -570,10 +581,10 @@ def extract_maintenance_groups_and_remaining( ) if is_start: - # 如果是第一个维修表,跳过(不加入 groups,也不标记 used_ids) + # 如果是第一个维修表,跳过(不加入 groups,也不标记 used_ids) if not skipped_first: skipped_first = True - # 注意:这里不清空 current_group,因为前面不可能有 group(第一个 start) + # 注意:这里不清空 current_group,因为前面不可能有 group(第一个 start) current_group = None # 确保不会把之前的非 start 内容误加 continue # 跳过这个 record,不加入任何 group @@ -586,7 +597,7 @@ def extract_maintenance_groups_and_remaining( if current_group is not None: current_group.append(record) - # 处理最后一个 group(仅当有有效 group 时) + # 处理最后一个 group(仅当有有效 group 时) if current_group is not None: groups.append(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): """ 通过语义判断表格是否为操作项目表(基于第一行内容),并分组。 - 注意:第一个识别到的操作项目表及其后续记录组成的 group 会被跳过(不加入 groups), + 注意:第一个识别到的操作项目表及其后续记录组成的 group 会被跳过(不加入 groups), 但这些记录仍保留在 remaining_records 中。 普通起始标记:组 = [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: # 决定是否保留当前积累的 group if not first_group_skipped: - # 跳过第一个 group:不清空 current_group,也不加入 groups 或 used_ids + # 跳过第一个 group:不清空 current_group,也不加入 groups 或 used_ids first_group_skipped = True # 注意:这里不把 current_group 加入 groups,也不更新 used_ids else: @@ -671,7 +682,7 @@ def extract_operation_groups_and_remaining(records, max_following_for_last=23): if current_group is not None: current_group.append(record) - # 处理最后一个 group(如果存在) + # 处理最后一个 group(如果存在) if current_group is not None: if not first_group_skipped: # 整个数据中只有一个 group,且是第一个 → 跳过它 @@ -680,7 +691,7 @@ def extract_operation_groups_and_remaining(records, max_following_for_last=23): groups.append(current_group) used_ids.update(r["id"] for r in current_group) - # —————— 后处理:截断最后一个 group(如果需要) —————— + # —————— 后处理:截断最后一个 group(如果需要) —————— if groups and max_following_for_last is not None: last_group = groups[-1] if len(last_group) > max_following_for_last + 1: # +1 是起始记录本身 @@ -729,7 +740,7 @@ def extract_wxanli(records): # 收集当前切片 slice_records = [] - # 添加前2个记录(如果存在且未超出边界) + # 添加前2个记录(如果存在且未超出边界) for offset in [2, 1]: prev_idx = start_pos - offset if prev_idx >= 0: @@ -810,85 +821,67 @@ def split_content_bbox(data, max_length=8000): -# def group_records(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 的标题 - """ +def _group_records_without_parent_context(records): result = [] current_group = [] - - # 维护标题栈:level -> {text, page_idx, bbox} - heading_stack = {} - + for i, ins in enumerate(records): - text_level = ins.get('text_level') - - 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 'text_level' in ins and ins['text_level']: + # 当遇到新的text_level时,如果current_group非空,则先将其添加到结果中 if current_group: result.append(current_group) - current_group = [] - - # 注入父级标题(如果当前层级 > 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] + current_group = [] # 开始新一组 + current_group.append(ins) else: - # 非标题记录,继续添加到当前组 + # 如果当前记录没有text_level或其值为空,则继续添加到当前组 current_group.append(ins) + # 别忘了将最后一个组添加到结果中 if current_group: result.append(current_group) 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): """ 处理没有 bbox 信息的文件数据。 @@ -1120,7 +1113,7 @@ async def data_replace(data, prefix): 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}, ...] """ chunks = [] @@ -1159,7 +1152,7 @@ def filter_content_bbox(data, max_length=3000): current_content = "" current_highlight_list = [] - # 如果单个 added_content 超过 max_length,仍需保留(避免丢弃) + # 如果单个 added_content 超过 max_length,仍需保留(避免丢弃) if len(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}") return None except json.JSONDecodeError: - print("错误:JSON 文件格式不正确") + print("错误:JSON 文件格式不正确") 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,则将其合并到下一个切片的开头 @@ -1358,110 +1351,6 @@ def find_ship_info_by_hull(json_file_path, data): ins['content'] = content[:newline_idx] + suffix + content[newline_idx:] 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'): """ 根据原始文件名,在指定目录及其子目录下查找对应的 _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_check = chunk_check(slices, 8000) slices_check = merge_short_slices(slices_check, min_length=30,filename=filename) + print(11111111111111111111111111111111111111) + print(len(images)) return {"slices": slices_check, "images": images} except json.JSONDecodeError: @@ -1596,7 +1487,7 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) - if __name__ == "__main__": - filepath ="/app/mineru_output/舰船抗沉损管训练仿真系统研究_content_list.json" + filepath = r"E:\ZKYNLP\Hjunproject\project0506\kgrag\舰船抗沉损管训练仿真系统研究_content_list.json" try: with open(filepath, 'r', encoding='utf-8') as file: data = json.load(file) @@ -1614,4 +1505,4 @@ if __name__ == "__main__": for ins in slices_2[:30]: - print(ins) + print(ins) \ No newline at end of file