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
40 changes: 18 additions & 22 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 @@ -204,18 +204,6 @@ def postprocess(
return (f1, recall, precision, vis_img) if visualize else (f1, recall, precision)


def _has_katex_error(img: np.ndarray) -> bool:
"""检测渲染图像中是否含有 KaTeX 红色报错文字(#cc0000 ≈ BGR(0,0,204))。"""
if img is None:
return True
b = img[:, :, 0].astype(np.int32)
g = img[:, :, 1].astype(np.int32)
r = img[:, :, 2].astype(np.int32)
# KaTeX error color: #cc0000 → RGB(204,0,0) → BGR(0,0,204)
error_mask = (np.abs(r - 204) < 30) & (g < 30) & (b < 30)
return int(error_mask.sum()) >= 5


class FastCDM:
def __init__(self, chromedriver: str = None) -> None:
self.chromedriver = chromedriver
Expand Down Expand Up @@ -291,17 +279,25 @@ 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]:
if self.render_worker is None:
return [RenderResult(None, True, "Renderer is unavailable", 0, 0) for _ in latex_list]
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 @@ -317,14 +313,14 @@ 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:
render_results = self.render_results([gt_latex, pred_latex])
if (
len(render_results) != 2
or any(result.error or result.image is None for result in render_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]

if _has_katex_error(gt_img) or _has_katex_error(pred_img):
self.render_failure_count += 1
gt_img, pred_img = render_results[0].image, render_results[1].image

result = postprocess(gt_img, pred_img, gt_color_map, pred_color_map, visualize)
return result
Expand Down
88 changes: 78 additions & 10 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,15 @@
from webdriver_manager.chrome import ChromeDriverManager


@dataclass
class RenderResult:
image: Optional[np.ndarray]
error: bool
error_text: Optional[str]
width: int
height: int


class RenderWorker:
"""
一个使用 Selenium Headless Chrome 渲染HTML内容的工具类。
Expand Down Expand Up @@ -76,15 +87,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, cv2.error) as exc:
return [RenderResult(None, True, type(exc).__name__, 0, 0) for _ in contents]

def _render(self, contents: List[str]) -> List[RenderResult]:
"""
渲染一组内容并返回每个元素的截图。
"""
# 通过JS调用页面内的render函数
self.driver.execute_script(
"document.body.classList.remove('rendering-complete');"
)
self.driver.execute_script(f"render({contents}, false)")
self.driver.execute_script("render(arguments[0], false)", contents)

# 等待JS渲染完成的信号
WebDriverWait(self.driver, self.timeout).until(
Expand All @@ -106,20 +125,56 @@ def render(self, contents: List[str]) -> List[np.ndarray]:
)
self.driver.set_window_size(self.window_fix_width, target_height)

dom_results = self.driver.execute_script(
"return [...document.querySelectorAll('.screenshot')].map(element => {"
" const rect = element.getBoundingClientRect();"
" const error = element.querySelector('.katex-error');"
" return {"
" error: Boolean(error || element.dataset.renderError),"
" errorText: element.dataset.renderError || (error ? error.textContent : null),"
" width: Math.ceil(Math.max(element.scrollWidth, rect.width)),"
" height: Math.ceil(Math.max(element.scrollHeight, rect.height))"
" };"
"});"
)
if len(dom_results) != len(contents):
return [
RenderResult(None, True, "DOM result count mismatch", 0, 0)
for _ in contents
]

# 获取整个页面的截图
png = self.driver.get_screenshot_as_png()
nparr = np.frombuffer(png, np.uint8)
fullpage_img = cv2.imdecode(nparr, cv2.IMREAD_COLOR)
if fullpage_img is None or fullpage_img.size == 0:
return [
RenderResult(
None,
True,
result.get("errorText") or "Empty browser screenshot",
int(result.get("width") or 0),
int(result.get("height") or 0),
)
for result in dom_results
]

# 获取每个渲染元素的边界框
rects = self.get_rects()
cropped_imgs = []
if len(rects) != len(contents):
return [RenderResult(None, True, "Capture rectangle count mismatch", 0, 0) for _ in contents]
results = []
img_h, img_w = fullpage_img.shape[:2]

# 根据边界框裁剪出每个元素的图像
for rect in rects:
if rect is None:
cropped_imgs.append(None)
for rect, dom_result in zip(rects, dom_results):
width = int(dom_result.get("width") or 0)
height = int(dom_result.get("height") or 0)
error_text = dom_result.get("errorText")
if rect is None or rect[2] <= 0 or rect[3] <= 0 or rect[0] < 0 or rect[1] < 0 or rect[0] + rect[2] > img_w or rect[1] + rect[3] > img_h:
results.append(
RenderResult(None, True, error_text or "Invalid capture rectangle", width, height)
)
else:
x, y, w, h = rect
# 计算一个小的随机边距,让截图更自然
Expand All @@ -132,9 +187,22 @@ 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)

return cropped_imgs
if cropped.size == 0:
results.append(
RenderResult(None, True, error_text or "Empty cropped image", width, height)
)
else:
results.append(
RenderResult(
cropped,
bool(dom_result.get("error")),
error_text,
width,
height,
)
)

return results

def get_rects(self) -> list:
"""
Expand Down
11 changes: 8 additions & 3 deletions fastcdm/render/templates/formula.html
Original file line number Diff line number Diff line change
Expand Up @@ -50,18 +50,23 @@
container.appendChild(div);
}

renderMathInElement(container, {
for (const element of container.children) {
renderMathInElement(element, {
delimiters: [
{ left: '$$', right: '$$', display: true },
{ left: '$', right: '$', display: false },
],
throwOnError: false,
throwOnError: true,
errorCallback: (message, error) => {
element.dataset.renderError = String(error || message);
},
maxExpand: 20000,
trust: true,
strict: false
});
}
document.body.classList.add('rendering-complete');
}
</script>
</body>
</html>
</html>