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 467cb866e0e0ad0a976f7d6766b04908bbea3836 Mon Sep 17 00:00:00 2001 From: liurui Date: Tue, 8 Sep 2026 15:19:04 +0800 Subject: [PATCH 2/2] fix(render): detect DOM render failures --- fastcdm/core.py | 26 ++++------ fastcdm/render/render_worker.py | 69 ++++++++++++++++++++++++--- fastcdm/render/templates/formula.html | 11 +++-- 3 files changed, 79 insertions(+), 27 deletions(-) diff --git a/fastcdm/core.py b/fastcdm/core.py index 7c1d56b..92e8b27 100644 --- a/fastcdm/core.py +++ b/fastcdm/core.py @@ -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 @@ -304,6 +292,8 @@ def render(self, latex_list: list) -> list: 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 ] @@ -323,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 diff --git a/fastcdm/render/render_worker.py b/fastcdm/render/render_worker.py index 3d35dff..10d8845 100644 --- a/fastcdm/render/render_worker.py +++ b/fastcdm/render/render_worker.py @@ -6,6 +6,7 @@ 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 @@ -21,7 +22,6 @@ class RenderResult: error_text: Optional[str] width: int height: int - error_type: Optional[str] = None class RenderWorker: @@ -88,6 +88,14 @@ def __init__(self, template_file: str, timeout: int = 15, driver_path: str = Non ) 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]: """ 渲染一组内容并返回每个元素的截图。 """ @@ -95,7 +103,7 @@ def render(self, contents: List[str]) -> List[RenderResult]: 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( @@ -117,20 +125,56 @@ def render(self, contents: List[str]) -> List[RenderResult]: ) 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() + 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: - results.append(RenderResult(None, True, "Invalid capture rectangle", 0, 0, "invalid_capture")) + 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 # 计算一个小的随机边距,让截图更自然 @@ -143,7 +187,20 @@ def render(self, contents: List[str]) -> List[RenderResult]: y2 = min(img_h, y + h + border_size) cropped = fullpage_img[y1:y2, x1:x2] - 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)) + 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 diff --git a/fastcdm/render/templates/formula.html b/fastcdm/render/templates/formula.html index 5fa85fc..d40a3f0 100644 --- a/fastcdm/render/templates/formula.html +++ b/fastcdm/render/templates/formula.html @@ -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'); } - \ No newline at end of file +