Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions fastcdm/__init__.py
Original file line number Diff line number Diff line change
@@ -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"
89 changes: 75 additions & 14 deletions fastcdm/core.py
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -291,17 +317,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:
"""
Expand All @@ -314,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

imgs = self.render([gt_latex, pred_latex])
if len(imgs) < 2 or imgs[0] is None or imgs[1] is None:
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))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里把所有预处理异常都固定标成 tokenize。process_for_katex() 的着色异常也会被错误归类,导致详细结果无法定位真实阶段。请拆分结构校验、tokenize 与 colorize 的异常边界。


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:
"""
Expand Down
23 changes: 17 additions & 6 deletions fastcdm/render/render_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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内容的工具类。
Expand Down Expand Up @@ -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]:
"""
渲染一组内容并返回每个元素的截图。
"""
Expand Down Expand Up @@ -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
# 计算一个小的随机边距,让截图更自然
Expand All @@ -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:
"""
Expand Down