From e590729669adb37c33e64083382e720ea2eb35c2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=B5=B5=E5=B0=8F=E8=92=99?= Date: Thu, 9 May 2024 17:25:24 +0800 Subject: [PATCH] fix span overlap by confidence,remove duplicate spans --- magic_pdf/libs/boxbase.py | 11 +++++++++++ magic_pdf/model/magic_model.py | 13 ++++++++++-- magic_pdf/pdf_parse_union_core.py | 5 ++++- magic_pdf/pre_proc/ocr_span_list_modify.py | 23 +++++++++++++++++++++- 4 files changed, 48 insertions(+), 4 deletions(-) diff --git a/magic_pdf/libs/boxbase.py b/magic_pdf/libs/boxbase.py index 25465d30..1e4e3430 100644 --- a/magic_pdf/libs/boxbase.py +++ b/magic_pdf/libs/boxbase.py @@ -161,6 +161,17 @@ def __is_overlaps_y_exceeds_threshold(bbox1, bbox2, overlap_ratio_threshold=0.8) def calculate_iou(bbox1, bbox2): + """ + 计算两个边界框的交并比(IOU)。 + + Args: + bbox1 (list[float]): 第一个边界框的坐标,格式为 [x1, y1, x2, y2],其中 (x1, y1) 为左上角坐标,(x2, y2) 为右下角坐标。 + bbox2 (list[float]): 第二个边界框的坐标,格式与 `bbox1` 相同。 + + Returns: + float: 两个边界框的交并比(IOU),取值范围为 [0, 1]。 + + """ # Determine the coordinates of the intersection rectangle x_left = max(bbox1[0], bbox2[0]) y_top = max(bbox1[1], bbox2[1]) diff --git a/magic_pdf/model/magic_model.py b/magic_pdf/model/magic_model.py index b6b96ba2..9285d202 100644 --- a/magic_pdf/model/magic_model.py +++ b/magic_pdf/model/magic_model.py @@ -448,6 +448,12 @@ class MagicModel: return text_spans def get_all_spans(self, page_no: int) -> list: + def remove_duplicate_spans(spans): + new_spans = [] + for span in spans: + if not any(span == existing_span for existing_span in new_spans): + new_spans.append(span) + return new_spans all_spans = [] model_page_info = self.__model_list[page_no] layout_dets = model_page_info["layout_dets"] @@ -461,7 +467,10 @@ class MagicModel: for layout_det in layout_dets: category_id = layout_det["category_id"] if category_id in allow_category_id_list: - span = {"bbox": layout_det["bbox"]} + span = { + "bbox": layout_det["bbox"], + "score": layout_det["score"] + } if category_id == 3: span["type"] = ContentType.Image elif category_id == 5: @@ -476,7 +485,7 @@ class MagicModel: span["content"] = layout_det["text"] span["type"] = ContentType.Text all_spans.append(span) - return all_spans + return remove_duplicate_spans(all_spans) def get_page_size(self, page_no: int): # 获取页面宽高 # 获取当前页的page对象 diff --git a/magic_pdf/pdf_parse_union_core.py b/magic_pdf/pdf_parse_union_core.py index fbd3b5ed..78a119e9 100644 --- a/magic_pdf/pdf_parse_union_core.py +++ b/magic_pdf/pdf_parse_union_core.py @@ -19,7 +19,8 @@ from magic_pdf.pre_proc.equations_replace import remove_chars_in_text_blocks, re from magic_pdf.pre_proc.ocr_detect_all_bboxes import ocr_prepare_bboxes_for_layout_split from magic_pdf.pre_proc.ocr_dict_merge import sort_blocks_by_layout, fill_spans_in_blocks, fix_block_spans, \ fix_discarded_block -from magic_pdf.pre_proc.ocr_span_list_modify import remove_overlaps_min_spans, get_qa_need_list_v2 +from magic_pdf.pre_proc.ocr_span_list_modify import remove_overlaps_min_spans, get_qa_need_list_v2, \ + remove_overlaps_low_confidence_spans from magic_pdf.pre_proc.resolve_bbox_conflict import check_useful_block_horizontal_overlap @@ -117,6 +118,8 @@ def parse_page_core(pdf_docs, magic_model, page_id, pdf_bytes_md5, imageWriter, else: raise Exception("parse_mode must be txt or ocr") + '''删除重叠spans中置信度较低的那些''' + spans, dropped_spans_by_confidence = remove_overlaps_low_confidence_spans(spans) '''删除重叠spans中较小的那些''' spans, dropped_spans_by_span_overlap = remove_overlaps_min_spans(spans) '''对image和table截图''' diff --git a/magic_pdf/pre_proc/ocr_span_list_modify.py b/magic_pdf/pre_proc/ocr_span_list_modify.py index 6dc6a8d5..9ed1ea2f 100644 --- a/magic_pdf/pre_proc/ocr_span_list_modify.py +++ b/magic_pdf/pre_proc/ocr_span_list_modify.py @@ -1,10 +1,31 @@ from loguru import logger from magic_pdf.libs.boxbase import calculate_overlap_area_in_bbox1_area_ratio, get_minbox_if_overlap_by_ratio, \ - __is_overlaps_y_exceeds_threshold + __is_overlaps_y_exceeds_threshold, calculate_iou from magic_pdf.libs.drop_tag import DropTag from magic_pdf.libs.ocr_content_type import ContentType, BlockType +def remove_overlaps_low_confidence_spans(spans): + dropped_spans = [] + # 删除重叠spans中置信度低的的那些 + for span1 in spans: + for span2 in spans: + if span1 != span2: + if calculate_iou(span1['bbox'], span2['bbox']) > 0.9: + if span1['score'] < span2['score']: + span_need_remove = span1 + else: + span_need_remove = span2 + if span_need_remove is not None and span_need_remove not in dropped_spans: + dropped_spans.append(span_need_remove) + + if len(dropped_spans) > 0: + for span_need_remove in dropped_spans: + spans.remove(span_need_remove) + span_need_remove['tag'] = DropTag.SPAN_OVERLAP + + return spans, dropped_spans + def remove_overlaps_min_spans(spans): dropped_spans = []