From 9c08223154c5cb82a53313f8dd1ff5a8e34bc16e Mon Sep 17 00:00:00 2001 From: liurui Date: Thu, 10 Sep 2026 15:23:38 +0800 Subject: [PATCH 1/2] refactor(render): add structured render result channel --- fastcdm/core.py | 14 ++++++++++---- fastcdm/render/render_worker.py | 23 +++++++++++++++++------ 2 files changed, 27 insertions(+), 10 deletions(-) diff --git a/fastcdm/core.py b/fastcdm/core.py index f986084..7c1d56b 100644 --- a/fastcdm/core.py +++ b/fastcdm/core.py @@ -1,4 +1,4 @@ -from fastcdm.render.render_worker import RenderWorker +from fastcdm.render.render_worker import RenderResult, RenderWorker from fastcdm.matcher import update_inliers, HungarianMatcher, SimpleAffineTransform from fastcdm.clean import ( clean, @@ -291,17 +291,23 @@ def render(self, latex_list: list) -> list: latex_strings = [ f"$${s}$$" if not s.startswith("$$") else s for s in latex_list ] - imgs = self.render_worker.render(latex_strings) + results = self.render_worker.render(latex_strings) except Exception as e: print("Rendering failed:") print("=" * 30) print(traceback.format_exc()) return [] - assert len(imgs) == len( + assert len(results) == len( latex_strings ), "Number of rendered images must match number of input strings" - return imgs + return [result.image for result in results] + + def render_results(self, latex_list: list) -> List[RenderResult]: + latex_strings = [ + f"$${s}$$" if not s.startswith("$$") else s for s in latex_list + ] + return self.render_worker.render(latex_strings) def compute(self, gt: str, pred: str, visualize: bool = False) -> tuple: """ diff --git a/fastcdm/render/render_worker.py b/fastcdm/render/render_worker.py index 111a586..3d35dff 100644 --- a/fastcdm/render/render_worker.py +++ b/fastcdm/render/render_worker.py @@ -2,7 +2,8 @@ import cv2 import random import numpy as np -from typing import List +from dataclasses import dataclass +from typing import List, Optional from selenium import webdriver from selenium.webdriver.chrome.service import Service as ChromeService @@ -13,6 +14,16 @@ from webdriver_manager.chrome import ChromeDriverManager +@dataclass +class RenderResult: + image: Optional[np.ndarray] + error: bool + error_text: Optional[str] + width: int + height: int + error_type: Optional[str] = None + + class RenderWorker: """ 一个使用 Selenium Headless Chrome 渲染HTML内容的工具类。 @@ -76,7 +87,7 @@ def __init__(self, template_file: str, timeout: int = 15, driver_path: str = Non EC.presence_of_all_elements_located((By.ID, "container")) ) - def render(self, contents: List[str]) -> List[np.ndarray]: + def render(self, contents: List[str]) -> List[RenderResult]: """ 渲染一组内容并返回每个元素的截图。 """ @@ -113,13 +124,13 @@ def render(self, contents: List[str]) -> List[np.ndarray]: # 获取每个渲染元素的边界框 rects = self.get_rects() - cropped_imgs = [] + results = [] img_h, img_w = fullpage_img.shape[:2] # 根据边界框裁剪出每个元素的图像 for rect in rects: if rect is None: - cropped_imgs.append(None) + results.append(RenderResult(None, True, "Invalid capture rectangle", 0, 0, "invalid_capture")) else: x, y, w, h = rect # 计算一个小的随机边距,让截图更自然 @@ -132,9 +143,9 @@ def render(self, contents: List[str]) -> List[np.ndarray]: y2 = min(img_h, y + h + border_size) cropped = fullpage_img[y1:y2, x1:x2] - cropped_imgs.append(cropped) + results.append(RenderResult(cropped, cropped.size == 0, "Empty cropped image" if cropped.size == 0 else None, w, h, "empty_image" if cropped.size == 0 else None)) - return cropped_imgs + return results def get_rects(self) -> list: """ From 3cf42735e94a0ca94fc26602eee65f745334994f Mon Sep 17 00:00:00 2001 From: liurui Date: Tue, 8 Sep 2026 15:19:10 +0800 Subject: [PATCH 2/2] feat(core): expose detailed failure results --- fastcdm/__init__.py | 4 +-- fastcdm/core.py | 75 +++++++++++++++++++++++++++++++++++++++------ 2 files changed, 67 insertions(+), 12 deletions(-) diff --git a/fastcdm/__init__.py b/fastcdm/__init__.py index f75a642..9a73c46 100644 --- a/fastcdm/__init__.py +++ b/fastcdm/__init__.py @@ -1,5 +1,5 @@ -from .core import FastCDM +from .core import CDMResult, FailureReport, FastCDM from .clean import clean -__all__ = ["FastCDM", "clean"] +__all__ = ["CDMResult", "FailureReport", "FastCDM", "clean"] __version__ = "0.1.4" diff --git a/fastcdm/core.py b/fastcdm/core.py index 7c1d56b..3bb34b7 100644 --- a/fastcdm/core.py +++ b/fastcdm/core.py @@ -11,7 +11,8 @@ import cv2 import numpy as np -from typing import List, Tuple +from dataclasses import dataclass +from typing import List, Literal, Optional, Tuple from pathlib import Path from skimage.measure import ransac import traceback @@ -23,6 +24,31 @@ TEMPLATE_FILE = root_dir / "render" / "templates" / "formula.html" +@dataclass +class FailureReport: + stage: Literal["tokenize", "colorize", "render", "capture", "postprocess"] + error_type: str + message: str + latex_summary: Optional[str] = None + attempt: int = 1 + renderer_rebuilt: bool = False + + +@dataclass +class CDMResult: + f1: Optional[float] + recall: Optional[float] + precision: Optional[float] + visualization: Optional[np.ndarray] + status: Literal["ok", "preprocess_failed", "render_failed", "postprocess_failed"] + failure: Optional[FailureReport] = None + + +def _latex_summary(gt: str, pred: str, limit: int = 160) -> str: + value = "GT: {} | Pred: {}".format((gt or "").strip(), (pred or "").strip()) + return value if len(value) <= limit else value[: limit - 1] + "…" + + def preprocess(s: str): # --- 第一步:清洗与分词 --- clean_s = clean(s) @@ -320,20 +346,49 @@ def compute(self, gt: str, pred: str, visualize: bool = False) -> tuple: 返回: tuple: 包含 F1 分数、召回率和准确率的元组。 """ - gt_latex, gt_color_map = preprocess(gt) - pred_latex, pred_color_map = preprocess(pred) + result = self.compute_detailed(gt, pred, visualize=visualize) + if result.status != "ok": + return (0, 0, 0, None) if visualize else (0, 0, 0) + if visualize: + return result.f1, result.recall, result.precision, result.visualization + return result.f1, result.recall, result.precision + + def compute_detailed(self, gt: str, pred: str, visualize: bool = False) -> CDMResult: + summary = _latex_summary(gt, pred) + try: + gt_latex, gt_color_map = preprocess(gt) + pred_latex, pred_color_map = preprocess(pred) + except Exception as exc: + return CDMResult(None, None, None, None, "preprocess_failed", FailureReport("tokenize", type(exc).__name__, str(exc), summary)) - imgs = self.render([gt_latex, pred_latex]) - if len(imgs) < 2 or imgs[0] is None or imgs[1] is None: + try: + render_results = self.render_results([gt_latex, pred_latex]) + except Exception as exc: self.render_failure_count += 1 - return (0, 0, 0, None) if visualize else (0, 0, 0) - gt_img, pred_img = imgs[0], imgs[1] + return CDMResult(None, None, None, None, "render_failed", FailureReport("render", type(exc).__name__, str(exc), summary)) - if _has_katex_error(gt_img) or _has_katex_error(pred_img): + if len(render_results) != 2: + self.render_failure_count += 1 + return CDMResult(None, None, None, None, "render_failed", FailureReport("render", "result_count_mismatch", "Expected two render results", summary)) + failed = next((item for item in render_results if item.error or item.image is None), None) + if failed is not None: self.render_failure_count += 1 + return CDMResult( + None, None, None, None, "render_failed", + FailureReport("render", "katex_error" if failed.error_text else "empty_image", failed.error_text or "Rendering produced no image", summary), + ) + + try: + metrics = postprocess(render_results[0].image, render_results[1].image, gt_color_map, pred_color_map, visualize) + except Exception as exc: + return CDMResult(None, None, None, None, "postprocess_failed", FailureReport("postprocess", type(exc).__name__, str(exc), summary)) - result = postprocess(gt_img, pred_img, gt_color_map, pred_color_map, visualize) - return result + if visualize: + f1, recall, precision, visualization = metrics + else: + f1, recall, precision = metrics + visualization = None + return CDMResult(float(f1), float(recall), float(precision), visualization, "ok") def batch_compute(self, gt_list: list, pred_list: list) -> list: """