diff --git a/mineru/backend/hybrid/hybrid_analyze.py b/mineru/backend/hybrid/hybrid_analyze.py index 617a1193..095aa2d4 100644 --- a/mineru/backend/hybrid/hybrid_analyze.py +++ b/mineru/backend/hybrid/hybrid_analyze.py @@ -23,8 +23,7 @@ from mineru.backend.pipeline.model_init import ( HybridModelSingleton, run_layout_inference, run_mfr_inference, - run_ocr_det_inference, - run_ocr_rec_inference, + run_ocr_inference, ) from mineru.backend.vlm.vlm_analyze import ( ModelSingleton, @@ -125,7 +124,7 @@ def ocr_det( page_mfd_res, useful_list ) bgr_image = cv2.cvtColor(new_image, cv2.COLOR_RGB2BGR) - ocr_res = run_ocr_det_inference( + ocr_res = run_ocr_inference( hybrid_pipeline_model.ocr_model.ocr, bgr_image, mfd_res=adjusted_mfdetrec_res, @@ -202,7 +201,7 @@ def ocr_det( # 批处理检测 det_batch_size = min(len(batch_images), batch_ratio * OCR_DET_BASE_BATCH_SIZE) - batch_results = run_ocr_det_inference( + batch_results = run_ocr_inference( hybrid_pipeline_model.ocr_model.text_detector.batch_predict, batch_images, det_batch_size, @@ -488,7 +487,7 @@ def _process_ocr_and_formulas( img_crop_list.append(ocr_res.pop('np_img')) if len(img_crop_list) > 0: # Process OCR - ocr_result_list = run_ocr_rec_inference( + ocr_result_list = run_ocr_inference( hybrid_pipeline_model.ocr_model.ocr, img_crop_list, det=False, diff --git a/mineru/backend/hybrid/hybrid_model_output_to_middle_json.py b/mineru/backend/hybrid/hybrid_model_output_to_middle_json.py index aa03e4cb..06821c64 100644 --- a/mineru/backend/hybrid/hybrid_model_output_to_middle_json.py +++ b/mineru/backend/hybrid/hybrid_model_output_to_middle_json.py @@ -14,12 +14,16 @@ from mineru.backend.utils.para_block_utils import ( ) from mineru.backend.hybrid.hybrid_magic_model import MagicModel from mineru.backend.utils.runtime_utils import cross_page_table_merge -from mineru.backend.pipeline.model_init import run_ocr_rec_inference +from mineru.backend.pipeline.model_init import run_ocr_inference from mineru.utils.config_reader import get_table_enable from mineru.utils.cut_image import cut_image_and_table from mineru.utils.enum_class import ContentType, BlockType from mineru.utils.hash_utils import bytes_md5 from mineru.utils.ocr_utils import OcrConfidence, rotate_vertical_crop_if_needed +from mineru.utils.span_pre_proc import ( + _clear_post_ocr_fallback, + _restore_post_ocr_fallback, +) from mineru.utils.title_level_postprocess import apply_title_leveling_to_pdf_info from mineru.utils.pdfium_guard import close_pdfium_child, close_pdfium_document, pdfium_guard from mineru.version import __version__ @@ -139,7 +143,7 @@ def _apply_post_ocr(pdf_info_list, hybrid_pipeline_model): img_crop_list.append(rotate_vertical_crop_if_needed(span['np_img'])) span.pop('np_img') if len(img_crop_list) > 0: - ocr_res_list = run_ocr_rec_inference( + ocr_res_list = run_ocr_inference( hybrid_pipeline_model.ocr_model.ocr, img_crop_list, det=False, @@ -152,6 +156,9 @@ def _apply_post_ocr(pdf_info_list, hybrid_pipeline_model): if ocr_score > OcrConfidence.min_confidence: span['content'] = ocr_text span['score'] = float(f"{ocr_score:.3f}") + _clear_post_ocr_fallback(span) + elif _restore_post_ocr_fallback(span): + continue else: span['content'] = '' span['score'] = 0.0 diff --git a/mineru/backend/pipeline/batch_analyze.py b/mineru/backend/pipeline/batch_analyze.py index 78f3f1a0..291a01da 100644 --- a/mineru/backend/pipeline/batch_analyze.py +++ b/mineru/backend/pipeline/batch_analyze.py @@ -13,8 +13,7 @@ from .model_init import ( AtomModelSingleton, run_layout_inference, run_mfr_inference, - run_ocr_det_inference, - run_ocr_rec_inference, + run_ocr_inference, ) from .model_list import AtomicModel from ...utils.config_reader import ( @@ -534,7 +533,7 @@ class BatchAnalyze: if inline_mask_boxes else bgr_image ) - ocr_result = run_ocr_det_inference( + ocr_result = run_ocr_inference( det_ocr_engine.ocr, det_image, rec=False )[0] if ocr_result and formula_mask_boxes: @@ -565,7 +564,7 @@ class BatchAnalyze: enable_merge_det_boxes=False, ) cropped_img_list = [item["cropped_img"] for item in rec_img_list] - ocr_res_list = run_ocr_rec_inference( + ocr_res_list = run_ocr_inference( ocr_engine.ocr, cropped_img_list, det=False, @@ -731,7 +730,7 @@ class BatchAnalyze: # 批处理检测 det_batch_size = min(len(batch_images), self.batch_ratio * OCR_DET_BASE_BATCH_SIZE) - batch_results = run_ocr_det_inference( + batch_results = run_ocr_inference( ocr_model.text_detector.batch_predict, batch_images, det_batch_size ) @@ -793,7 +792,7 @@ class BatchAnalyze: bgr_image, adjusted_mfdetrec_res, ) - ocr_res = run_ocr_det_inference( + ocr_res = run_ocr_inference( ocr_model.ocr, det_image, mfd_res=adjusted_mfdetrec_res, @@ -852,7 +851,7 @@ class BatchAnalyze: atom_model_name=AtomicModel.OCR, lang=lang ) - ocr_res_list = run_ocr_rec_inference( + ocr_res_list = run_ocr_inference( ocr_model.ocr, img_crop_list, det=False, tqdm_enable=True )[0] @@ -923,7 +922,7 @@ class BatchAnalyze: ) seal_crop_bgr = cv2.cvtColor(seal_crop_rgb, cv2.COLOR_RGB2BGR) - seal_ocr_res = run_ocr_det_inference( + seal_ocr_res = run_ocr_inference( seal_ocr_model.ocr, seal_crop_bgr, det=True, rec=True )[0] if not seal_ocr_res: diff --git a/mineru/backend/pipeline/model_init.py b/mineru/backend/pipeline/model_init.py index 12738f46..9b20c810 100644 --- a/mineru/backend/pipeline/model_init.py +++ b/mineru/backend/pipeline/model_init.py @@ -22,8 +22,7 @@ PIPELINE_MODEL_INIT_LOCK = threading.RLock() # 这些锁保护 pipeline 与 hybrid 共享的 atom model/native 模型推理调用,避免多线程同时进入同一个模型对象。 PIPELINE_LAYOUT_INFERENCE_LOCK = threading.RLock() PIPELINE_MFR_INFERENCE_LOCK = threading.RLock() -PIPELINE_OCR_DET_INFERENCE_LOCK = threading.RLock() -PIPELINE_OCR_REC_INFERENCE_LOCK = threading.RLock() +PIPELINE_OCR_INFERENCE_LOCK = threading.RLock() # 临时关闭 pipeline/hybrid 共享推理阶段锁;需要回滚实验时可通过环境变量重新打开。 PIPELINE_INFERENCE_LOCKS_ENABLED = os.getenv( 'MINERU_ENABLE_PIPELINE_INFERENCE_LOCKS', 'False' @@ -53,17 +52,10 @@ def run_mfr_inference(inference_callable, *args, **kwargs): ) -def run_ocr_det_inference(inference_callable, *args, **kwargs): - """按实验开关执行共享 OCR det 模型调用。""" +def run_ocr_inference(inference_callable, *args, **kwargs): + """按实验开关执行共享 OCR native 模型调用。""" return _run_with_inference_lock( - PIPELINE_OCR_DET_INFERENCE_LOCK, inference_callable, *args, **kwargs - ) - - -def run_ocr_rec_inference(inference_callable, *args, **kwargs): - """按实验开关执行共享 OCR rec 模型调用。""" - return _run_with_inference_lock( - PIPELINE_OCR_REC_INFERENCE_LOCK, inference_callable, *args, **kwargs + PIPELINE_OCR_INFERENCE_LOCK, inference_callable, *args, **kwargs ) MFR_MODEL = os.getenv('MINERU_FORMULA_CH_SUPPORT', 'False') diff --git a/mineru/backend/pipeline/model_json_to_middle_json.py b/mineru/backend/pipeline/model_json_to_middle_json.py index 7ac043ff..5c4cc9b9 100644 --- a/mineru/backend/pipeline/model_json_to_middle_json.py +++ b/mineru/backend/pipeline/model_json_to_middle_json.py @@ -7,7 +7,7 @@ from mineru.backend.utils.html_image_utils import replace_inline_table_images from mineru.backend.utils.runtime_utils import cross_page_table_merge from mineru.backend.pipeline.model_init import ( AtomModelSingleton, - run_ocr_rec_inference, + run_ocr_inference, ) from mineru.backend.pipeline.para_split import para_split from mineru.utils.char_utils import full_to_half @@ -19,6 +19,10 @@ from mineru.utils.ocr_utils import OcrConfidence, rotate_vertical_crop_if_needed from mineru.version import __version__ from mineru.utils.hash_utils import bytes_md5 from mineru.utils.pdfium_guard import close_pdfium_child, close_pdfium_document, pdfium_guard +from mineru.utils.span_pre_proc import ( + _clear_post_ocr_fallback, + _restore_post_ocr_fallback, +) def page_model_info_to_page_info(page_model_info, image_dict, page, image_writer, page_index, ocr_enable=False): @@ -235,7 +239,7 @@ def _apply_post_ocr(pdf_info_list, lang=None): det_db_box_thresh=0.3, lang=lang ) - ocr_res_list = run_ocr_rec_inference( + ocr_res_list = run_ocr_inference( ocr_model.ocr, img_crop_list, det=False, tqdm_enable=True )[0] assert len(ocr_res_list) == len( @@ -245,6 +249,9 @@ def _apply_post_ocr(pdf_info_list, lang=None): if ocr_score > OcrConfidence.min_confidence: span['content'] = ocr_text span['score'] = float(f"{ocr_score:.3f}") + _clear_post_ocr_fallback(span) + elif _restore_post_ocr_fallback(span): + continue else: span['content'] = '' span['score'] = 0.0 diff --git a/mineru/utils/span_pre_proc.py b/mineru/utils/span_pre_proc.py index 7648d815..cdb77694 100644 --- a/mineru/utils/span_pre_proc.py +++ b/mineru/utils/span_pre_proc.py @@ -15,6 +15,15 @@ from mineru.utils.pdf_text_tool import get_lines_from_chars, get_page_chars from mineru.utils.pdfium_guard import close_pdfium_child, pdfium_guard MAX_NATIVE_TEXT_CHARS_PER_PAGE = 65535 +PRIVATE_USE_AREA_START = 0xE000 +PRIVATE_USE_AREA_END = 0xF8FF +PRIVATE_USE_TEXT_COUNT_THRESHOLD = 2 +PRIVATE_USE_TEXT_RATIO_THRESHOLD = 0.05 +PRIVATE_USE_TEXT_RUN_THRESHOLD = 2 +POST_OCR_FALLBACK_CONTENT_KEY = '_post_ocr_fallback_content' +POST_OCR_FALLBACK_SCORE_KEY = '_post_ocr_fallback_score' +POST_OCR_REASON_KEY = '_post_ocr_reason' +POST_OCR_REASON_PRIVATE_USE_TEXT = 'private_use_text' def __replace_ligatures(text: str): @@ -141,6 +150,8 @@ def _prepare_post_ocr_spans(need_ocr_spans, spans, pil_img, scale): span_img = cv2.cvtColor(np.array(span_pil_img), cv2.COLOR_RGB2BGR) # 计算span的对比度,低于0.17的span不进行ocr,等于0.17的临界框保留给后置OCR。 if calculate_contrast(span_img, img_mode='bgr') < 0.17: + if _restore_post_ocr_fallback(span): + continue if span in spans: spans.remove(span) continue @@ -267,9 +278,18 @@ def fill_char_in_spans(spans, all_chars, median_span_height): need_ocr_spans = [] for span in spans: + private_use_signal = _get_private_use_text_signal(span['chars']) + should_post_ocr_private_use = _should_fallback_to_post_ocr_for_private_use_text( + private_use_signal + ) chars_to_content(span) + if should_post_ocr_private_use and span.get('content'): + span[POST_OCR_FALLBACK_CONTENT_KEY] = span['content'] + span[POST_OCR_FALLBACK_SCORE_KEY] = span.get('score', 1.0) + span[POST_OCR_REASON_KEY] = POST_OCR_REASON_PRIVATE_USE_TEXT + need_ocr_spans.append(span) # 有的span中虽然没有字但有一两个空的占位符,用宽高和content长度过滤 - if len(span['content']) * span['height'] < span['width'] * 0.5: + elif len(span['content']) * span['height'] < span['width'] * 0.5: # logger.info(f"maybe empty span: {len(span['content'])}, {span['height']}, {span['width']}") need_ocr_spans.append(span) del span['height'], span['width'] @@ -280,6 +300,83 @@ LINE_STOP_FLAG = ('.', '!', '?', '。', '!', '?', ')', ')', '"', '”', ': LINE_START_FLAG = ('(', '(', '"', '“', '【', '{', '《', '<', '「', '『', '【', '[',) Span_Height_Ratio = 0.33 # 字符的中轴和span的中轴高度差不能超过1/3span高度 +SCRIPT_BODY_HEIGHT_RATIO = 0.9 +SCRIPT_CENTER_TOLERANCE_RATIO = 0.12 + + +def _is_private_use_char(char: str) -> bool: + """判断单个字符是否落在 Unicode 私用区,用于识别字体映射异常。""" + return ( + len(char) == 1 + and PRIVATE_USE_AREA_START <= ord(char) <= PRIVATE_USE_AREA_END + ) + + +def _get_private_use_text_signal(chars): + """统计 span 字符中的私用区信号,供局部后置 OCR 决策使用。""" + pua_count = 0 + text_char_count = 0 + current_pua_run = 0 + max_pua_run = 0 + + for char in chars: + for text_char in char.get('char', ''): + if text_char.isspace(): + current_pua_run = 0 + continue + + text_char_count += 1 + if _is_private_use_char(text_char): + pua_count += 1 + current_pua_run += 1 + max_pua_run = max(max_pua_run, current_pua_run) + else: + current_pua_run = 0 + + pua_ratio = 0.0 + if text_char_count > 0: + pua_ratio = pua_count / text_char_count + + return { + 'pua_count': pua_count, + 'text_char_count': text_char_count, + 'pua_ratio': pua_ratio, + 'max_pua_run': max_pua_run, + } + + +def _should_fallback_to_post_ocr_for_private_use_text(signal) -> bool: + """连续或高占比 PUA 才转后置 OCR,降低孤立私用符号误召回。""" + pua_count = signal['pua_count'] + if pua_count < PRIVATE_USE_TEXT_COUNT_THRESHOLD: + return False + + return ( + signal['max_pua_run'] >= PRIVATE_USE_TEXT_RUN_THRESHOLD + or signal['pua_ratio'] >= PRIVATE_USE_TEXT_RATIO_THRESHOLD + ) + + +def _clear_post_ocr_fallback(span): + """清理后置 OCR 内部兜底字段,避免进入最终 middle-json 输出。""" + span.pop(POST_OCR_FALLBACK_CONTENT_KEY, None) + span.pop(POST_OCR_FALLBACK_SCORE_KEY, None) + span.pop(POST_OCR_REASON_KEY, None) + + +def _restore_post_ocr_fallback(span) -> bool: + """在后置 OCR 无法使用时恢复原始文本兜底,返回是否已恢复。""" + if POST_OCR_FALLBACK_CONTENT_KEY not in span: + _clear_post_ocr_fallback(span) + return False + + span['content'] = span[POST_OCR_FALLBACK_CONTENT_KEY] + if POST_OCR_FALLBACK_SCORE_KEY in span: + span['score'] = span[POST_OCR_FALLBACK_SCORE_KEY] + _clear_post_ocr_fallback(span) + return True + + def calculate_char_in_span(char_bbox, span_bbox, char, span_height_ratio=Span_Height_Ratio): char_center_x = (char_bbox[0] + char_bbox[2]) / 2 char_center_y = (char_bbox[1] + char_bbox[3]) / 2 @@ -315,6 +412,114 @@ def calculate_char_in_span(char_bbox, span_bbox, char, span_height_ratio=Span_He return False +def _get_char_bbox_metrics(char): + """提取字符 bbox 的宽高和中心点,统一兼容 list 与 pdftext Bbox 对象。""" + bbox = char['bbox'] + x0, y0, x1, y1 = [float(v) for v in bbox] + return { + 'width': x1 - x0, + 'height': y1 - y0, + 'center_y': (y0 + y1) / 2, + } + + +def _get_char_bbox_metrics_list(chars): + """预计算 span 内全部字符的 bbox 指标,避免上下标判断重复解析 bbox。""" + return [_get_char_bbox_metrics(char) for char in chars] + + +def _is_valid_script_reference_char(char, metrics) -> bool: + """过滤空白和退化 bbox,只用真实可见字符估计正文主带。""" + if char['char'] in {' ', '\r', '\n'}: + return False + + return metrics['height'] > 1 and metrics['width'] > 0 + + +def _get_body_axis(chars, char_metrics): + """根据同一 span 内最大高度字符簇估计正文中心线和正文高度。""" + valid_metrics = [ + metrics for char, metrics in zip(chars, char_metrics) + if _is_valid_script_reference_char(char, metrics) + ] + if not valid_metrics: + return None + + max_height = max(metrics['height'] for metrics in valid_metrics) + body_metrics = [ + metrics for metrics in valid_metrics + if metrics['height'] >= max_height * SCRIPT_BODY_HEIGHT_RATIO + ] + if not body_metrics: + return None + + return { + 'center_y': statistics.median(metrics['center_y'] for metrics in body_metrics), + 'height': statistics.median(metrics['height'] for metrics in body_metrics), + } + + +def _classify_char_script_roles(chars, char_metrics): + """按正文主带判断每个字符属于正文、上标或下标。""" + body_axis = _get_body_axis(chars, char_metrics) + if body_axis is None or body_axis['height'] <= 0: + return ['body'] * len(chars) + + tolerance = body_axis['height'] * SCRIPT_CENTER_TOLERANCE_RATIO + roles = [] + for char, metrics in zip(chars, char_metrics): + if not _is_valid_script_reference_char(char, metrics): + roles.append('body') + continue + + char_center_y = metrics['center_y'] + if char_center_y < body_axis['center_y'] - tolerance: + roles.append('sup') + elif char_center_y > body_axis['center_y'] + tolerance: + roles.append('sub') + else: + roles.append('body') + return roles + + +def _append_script_wrapped_text(parts, role, text): + """把连续同类上下标文本包裹成 HTML 标签,正文保持原样。""" + if not text: + return + if role == 'sup': + parts.append(f'{text}') + elif role == 'sub': + parts.append(f'{text}') + else: + parts.append(text) + + +def _wrap_script_runs(role_text_parts): + """合并连续正文、上标、下标 run,避免每个字符单独生成标签。""" + wrapped_parts = [] + current_role = None + current_text_parts = [] + + for role, text in role_text_parts: + if role != current_role: + _append_script_wrapped_text( + wrapped_parts, + current_role, + ''.join(current_text_parts), + ) + current_role = role + current_text_parts = [text] + else: + current_text_parts.append(text) + + _append_script_wrapped_text( + wrapped_parts, + current_role, + ''.join(current_text_parts), + ) + return ''.join(wrapped_parts) + + def chars_to_content(span): # 检查span中的char是否为空 if len(span['chars']) != 0: @@ -326,28 +531,31 @@ def chars_to_content(span): ): chars = sorted(chars, key=lambda x: x['char_idx']) + char_metrics = _get_char_bbox_metrics_list(chars) # Calculate the width of each character - char_widths = [char['bbox'][2] - char['bbox'][0] for char in chars] + char_widths = [metrics['width'] for metrics in char_metrics] # Calculate the median width median_width = statistics.median(char_widths) + script_roles = _classify_char_script_roles(chars, char_metrics) - parts = [] + role_text_parts = [] for idx, char1 in enumerate(chars): char2 = chars[idx + 1] if idx + 1 < len(chars) else None + role1 = script_roles[idx] + role2 = script_roles[idx + 1] if char2 else None # 如果下一个char的x0和上一个char的x1距离超过0.25个字符宽度,则需要在中间插入一个空格 + role_text_parts.append((role1, char1['char'])) if ( char2 and char2['bbox'][0] - char1['bbox'][2] > median_width * 0.25 and char1['char'] != ' ' and char2['char'] != ' ' ): - parts.append(char1['char']) - parts.append(' ') - else: - parts.append(char1['char']) + space_role = role1 if role1 == role2 else 'body' + role_text_parts.append((space_role, ' ')) - content = ''.join(parts) + content = _wrap_script_runs(role_text_parts) content = __replace_unicode(content) content = __replace_ligatures(content) span['content'] = content.strip() diff --git a/pyproject.toml b/pyproject.toml index b8b05b91..ed6fd3e0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -180,3 +180,10 @@ exclude_also = [ 'class .*\bProtocol\):', '@(abc\.)?abstractmethod', ] + +[tool.ruff] +line-length = 128 + +[tool.ruff.lint] +select = ["C", "E", "F", "W", "ANN"] +ignore = ["C901", "ANN204", "ANN401"]