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 f6c505b1840b7a060af05475c13d066cf55ccdbd Mon Sep 17 00:00:00 2001 From: liurui Date: Thu, 10 Sep 2026 15:26:32 +0800 Subject: [PATCH 2/2] fix(render): size viewport for long formulas --- fastcdm/render/render_worker.py | 42 +++++++++++++++++++++++++++++---- 1 file changed, 37 insertions(+), 5 deletions(-) diff --git a/fastcdm/render/render_worker.py b/fastcdm/render/render_worker.py index 3d35dff..7be42fd 100644 --- a/fastcdm/render/render_worker.py +++ b/fastcdm/render/render_worker.py @@ -24,6 +24,10 @@ class RenderResult: error_type: Optional[str] = None +def clamp_width(content_width: int, min_width: int, max_width: int, margin: int) -> int: + return min(max(content_width + margin, min_width), max_width) + + class RenderWorker: """ 一个使用 Selenium Headless Chrome 渲染HTML内容的工具类。 @@ -31,7 +35,15 @@ class RenderWorker: 并截取渲染后各元素的图像。 """ - def __init__(self, template_file: str, timeout: int = 15, driver_path: str = None): + def __init__( + self, + template_file: str, + timeout: int = 15, + driver_path: str = None, + min_width: int = 2000, + max_width: int = 16000, + horizontal_margin: int = 40, + ): # --- 配置浏览器选项 --- opts = Options() opts.add_argument("--headless") @@ -71,8 +83,12 @@ def __init__(self, template_file: str, timeout: int = 15, driver_path: str = Non self.timeout = timeout - # 定义一个固定的窗口宽度 - self.window_fix_width = 2000 + if min_width <= 0 or max_width < min_width or horizontal_margin < 0: + raise ValueError("Invalid render width configuration") + self.min_width = min_width + self.max_width = max_width + self.horizontal_margin = horizontal_margin + self.window_fix_width = min_width self.window_init_height = 300 self.driver.set_window_size(self.window_fix_width, self.window_init_height) @@ -102,6 +118,23 @@ def render(self, contents: List[str]) -> List[RenderResult]: EC.presence_of_element_located((By.CLASS_NAME, "rendering-complete")) ) + measured_widths = self.driver.execute_script( + "return [...document.querySelectorAll('.screenshot')]" + ".map(e => Math.ceil(Math.max(e.scrollWidth, e.getBoundingClientRect().width)));" + ) + if len(measured_widths) != len(contents): + return [RenderResult(None, True, "DOM result count mismatch", 0, 0) for _ in contents] + required_width = max([int(value or 0) for value in measured_widths] or [0]) + if required_width + self.horizontal_margin > self.max_width: + return [ + RenderResult(None, True, "Formula width exceeds maximum", int(width or 0), 0) + for width in measured_widths + ] + self.window_fix_width = clamp_width( + required_width, self.min_width, self.max_width, self.horizontal_margin + ) + self.driver.set_window_size(self.window_fix_width, self.window_init_height) + # 根据内容的总高度调整窗口大小,以确保能截取完整图像 scroll_height = self.driver.execute_script( "return document.getElementById('container').scrollHeight" @@ -164,8 +197,7 @@ def get_rects(self) -> list: w = int(size["width"]) h = int(size["height"]) - # 如果元素宽度超过窗口,可能是一个渲染错误,标记为None - if w > self.window_fix_width: + if w <= 0 or h <= 0 or x < 0 or y < 0 or x + w > self.window_fix_width: rects.append(None) else: rects.append((x, y, w, h))