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
14 changes: 10 additions & 4 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 @@ -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:
"""
Expand Down
65 changes: 54 additions & 11 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,14 +14,36 @@
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


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内容的工具类。
它可以加载一个HTML模板,通过JavaScript渲染内容(如数学公式),
并截取渲染后各元素的图像。
"""

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")
Expand Down Expand Up @@ -60,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)
Expand All @@ -76,7 +103,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 All @@ -91,6 +118,23 @@ def render(self, contents: List[str]) -> List[np.ndarray]:
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"
Expand All @@ -113,13 +157,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 +176,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 All @@ -153,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))
Expand Down