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
62 changes: 55 additions & 7 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 Down Expand Up @@ -217,6 +217,8 @@ def _has_katex_error(img: np.ndarray) -> bool:


class FastCDM:
RETRYABLE_RENDER_ERRORS = {"empty_image", "invalid_capture", "webdriver_error"}

def __init__(self, chromedriver: str = None) -> None:
self.chromedriver = chromedriver
self.render_failure_count: int = 0
Expand Down Expand Up @@ -291,17 +293,57 @@ 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 _recover_renderer(self, failed_attempt: int) -> bool:
if failed_attempt == 1 and self.render_worker is not None:
self.render_worker.driver.refresh()
return False
self.close()
self.init_render_worker()
return True

def _render_with_retries(self, latex_list: list):
last_results = []
last_exception = None
rebuilt = False
for attempt in range(1, 4):
try:
last_results = self.render_results(latex_list)
last_exception = None
except Exception as exc:
last_results = []
last_exception = exc

if last_exception is None and len(last_results) == len(latex_list):
failures = [item for item in last_results if item.error or item.image is None]
if not failures:
return last_results, attempt, rebuilt, None
if any(item.error_type not in self.RETRYABLE_RENDER_ERRORS for item in failures):
return last_results, attempt, rebuilt, None

if attempt < 3:
try:
rebuilt = self._recover_renderer(attempt) or rebuilt
except Exception as exc:
last_exception = exc
return last_results, 3, rebuilt, last_exception

def compute(self, gt: str, pred: str, visualize: bool = False) -> tuple:
"""
Expand All @@ -317,11 +359,17 @@ def compute(self, gt: str, pred: str, visualize: bool = False) -> tuple:
gt_latex, gt_color_map = preprocess(gt)
pred_latex, pred_color_map = preprocess(pred)

imgs = self.render([gt_latex, pred_latex])
if len(imgs) < 2 or imgs[0] is None or imgs[1] is None:
results, _, _, render_exception = self._render_with_retries(
[gt_latex, pred_latex]
)
if (
render_exception is not None
or len(results) < 2
or any(result.error or result.image is None for result in results[:2])
):
self.render_failure_count += 1
return (0, 0, 0, None) if visualize else (0, 0, 0)
gt_img, pred_img = imgs[0], imgs[1]
gt_img, pred_img = results[0].image, results[1].image

if _has_katex_error(gt_img) or _has_katex_error(pred_img):
self.render_failure_count += 1
Expand Down
40 changes: 34 additions & 6 deletions fastcdm/render/render_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,11 @@
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.common.exceptions import WebDriverException
from selenium.webdriver.chrome.service import Service as ChromeService
from selenium.webdriver.chrome.options import Options
from selenium.webdriver.common.by import By
Expand All @@ -13,6 +15,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 +88,23 @@ 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]:
if not contents:
return []
try:
return self._render(contents)
except WebDriverException as exc:
return [
RenderResult(None, True, f"{type(exc).__name__}: {exc}", 0, 0, "webdriver_error")
for _ in contents
]
except cv2.error as exc:
return [
RenderResult(None, True, f"{type(exc).__name__}: {exc}", 0, 0, "invalid_capture")
for _ in contents
]

def _render(self, contents: List[str]) -> List[RenderResult]:
"""
渲染一组内容并返回每个元素的截图。
"""
Expand Down Expand Up @@ -113,13 +141,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 +160,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