OCR-to-RAG pipeline for life science screenshots
PaddleOCR with preprocessing, scispaCy NER, figure/table detection, citation extraction, and chunked Markdown output with frontmatter. Includes watch mode and notebook reprocessing.
This commit is contained in:
15
src/ocr_pipeline/detectors/__init__.py
Normal file
15
src/ocr_pipeline/detectors/__init__.py
Normal file
@@ -0,0 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from .figures import CaptionExtractor, DetectedFigure, FigureDetector, detect_figures
|
||||
from .tables import DetectedTable, TableDetector, TableStructureRecognizer, detect_tables
|
||||
|
||||
__all__ = [
|
||||
"detect_figures",
|
||||
"FigureDetector",
|
||||
"CaptionExtractor",
|
||||
"DetectedFigure",
|
||||
"detect_tables",
|
||||
"TableDetector",
|
||||
"TableStructureRecognizer",
|
||||
"DetectedTable",
|
||||
]
|
||||
BIN
src/ocr_pipeline/detectors/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/detectors/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/detectors/__pycache__/figures.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/detectors/__pycache__/figures.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/detectors/__pycache__/tables.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/detectors/__pycache__/tables.cpython-313.pyc
Normal file
Binary file not shown.
166
src/ocr_pipeline/detectors/figures.py
Normal file
166
src/ocr_pipeline/detectors/figures.py
Normal file
@@ -0,0 +1,166 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DetectedFigure:
|
||||
bbox: list[float]
|
||||
confidence: float
|
||||
label: str
|
||||
caption: str | None = None
|
||||
image_crop: np.ndarray | None = None
|
||||
|
||||
|
||||
class FigureDetector:
|
||||
def __init__(self):
|
||||
self._predictor = None
|
||||
self._initialized = False
|
||||
|
||||
def _init(self):
|
||||
if self._initialized:
|
||||
return
|
||||
if not settings.detectors.figures.enabled:
|
||||
self._initialized = True
|
||||
return
|
||||
try:
|
||||
import layoutparser as lp
|
||||
|
||||
self._predictor = lp.Detectron2LayoutModel(
|
||||
config_path=settings.detectors.figures.model,
|
||||
label_map={0: "Text", 1: "Title", 2: "List", 3: "Table", 4: "Figure"},
|
||||
extra_config=[
|
||||
"MODEL.ROI_HEADS.SCORE_THRESH_TEST",
|
||||
settings.detectors.figures.confidence_threshold,
|
||||
],
|
||||
)
|
||||
self._initialized = True
|
||||
logger.info("figure_detector_initialized", model=settings.detectors.figures.model)
|
||||
except Exception as e:
|
||||
logger.warning("figure_detector_init_failed", error=str(e))
|
||||
self._initialized = True
|
||||
|
||||
def detect(self, image: np.ndarray) -> list[DetectedFigure]:
|
||||
if not settings.detectors.figures.enabled:
|
||||
return []
|
||||
|
||||
self._init()
|
||||
if self._predictor is None:
|
||||
return []
|
||||
|
||||
try:
|
||||
layout = self._predictor.detect(image)
|
||||
figures = []
|
||||
|
||||
for block in layout:
|
||||
if block.type == "Figure":
|
||||
x1, y1, x2, y2 = (
|
||||
block.block.x_1,
|
||||
block.block.y_1,
|
||||
block.block.x_2,
|
||||
block.block.y_2,
|
||||
)
|
||||
crop = image[int(y1) : int(y2), int(x1) : int(x2)]
|
||||
figures.append(
|
||||
DetectedFigure(
|
||||
bbox=[float(x1), float(y1), float(x2), float(y2)],
|
||||
confidence=float(block.score),
|
||||
label="Figure",
|
||||
image_crop=crop if crop.size > 0 else None,
|
||||
)
|
||||
)
|
||||
|
||||
logger.debug("figures_detected", count=len(figures))
|
||||
return figures
|
||||
except Exception as e:
|
||||
logger.warning("figure_detection_failed", error=str(e))
|
||||
return []
|
||||
|
||||
|
||||
class CaptionExtractor:
|
||||
def __init__(self):
|
||||
self._predictor = None
|
||||
self._initialized = False
|
||||
|
||||
def _init(self):
|
||||
if self._initialized:
|
||||
return
|
||||
if not settings.detectors.captions.enabled:
|
||||
self._initialized = True
|
||||
return
|
||||
try:
|
||||
import layoutparser as lp
|
||||
|
||||
self._predictor = lp.Detectron2LayoutModel(
|
||||
config_path=settings.detectors.captions.model,
|
||||
label_map={0: "Text", 1: "Title", 2: "List", 3: "Table", 4: "Figure"},
|
||||
extra_config=[
|
||||
"MODEL.ROI_HEADS.SCORE_THRESH_TEST",
|
||||
settings.detectors.captions.confidence_threshold,
|
||||
],
|
||||
)
|
||||
self._initialized = True
|
||||
except Exception as e:
|
||||
logger.warning("caption_detector_init_failed", error=str(e))
|
||||
self._initialized = True
|
||||
|
||||
def extract_captions(
|
||||
self, image: np.ndarray, figures: list[DetectedFigure]
|
||||
) -> list[DetectedFigure]:
|
||||
if not settings.detectors.captions.enabled or not figures:
|
||||
return figures
|
||||
|
||||
self._init()
|
||||
if self._predictor is None:
|
||||
return figures
|
||||
|
||||
try:
|
||||
layout = self._predictor.detect(image)
|
||||
captions = [
|
||||
b
|
||||
for b in layout
|
||||
if b.type in ("Title", "Text")
|
||||
and b.score > settings.detectors.captions.confidence_threshold
|
||||
]
|
||||
|
||||
for fig in figures:
|
||||
fig_x1, fig_y1, fig_x2, fig_y2 = fig.bbox
|
||||
fig_center_y = (fig_y1 + fig_y2) / 2
|
||||
|
||||
best_caption = None
|
||||
min_dist = float("inf")
|
||||
|
||||
for cap in captions:
|
||||
cap_y1, cap_y2 = cap.block.y_1, cap.block.y_2
|
||||
cap_center_y = (cap_y1 + cap_y2) / 2
|
||||
dist = abs(cap_center_y - fig_center_y)
|
||||
|
||||
if cap_y2 < fig_y1 or cap_y1 > fig_y2:
|
||||
if dist < min_dist:
|
||||
min_dist = dist
|
||||
best_caption = cap
|
||||
|
||||
if best_caption and min_dist < 200:
|
||||
fig.caption = best_caption.block.text
|
||||
|
||||
return figures
|
||||
except Exception as e:
|
||||
logger.warning("caption_extraction_failed", error=str(e))
|
||||
return figures
|
||||
|
||||
|
||||
def detect_figures(image: np.ndarray) -> list[DetectedFigure]:
|
||||
detector = FigureDetector()
|
||||
figures = detector.detect(image)
|
||||
|
||||
captioner = CaptionExtractor()
|
||||
figures = captioner.extract_captions(image, figures)
|
||||
|
||||
return figures
|
||||
202
src/ocr_pipeline/detectors/tables.py
Normal file
202
src/ocr_pipeline/detectors/tables.py
Normal file
@@ -0,0 +1,202 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DetectedTable:
|
||||
bbox: list[float]
|
||||
confidence: float
|
||||
label: str
|
||||
structure: list[dict[str, Any]] | None = None
|
||||
markdown: str | None = None
|
||||
image_crop: np.ndarray | None = None
|
||||
|
||||
|
||||
class TableDetector:
|
||||
def __init__(self):
|
||||
self._model = None
|
||||
self._processor = None
|
||||
self._initialized = False
|
||||
|
||||
def _init(self):
|
||||
if self._initialized:
|
||||
return
|
||||
if not settings.detectors.tables.enabled:
|
||||
self._initialized = True
|
||||
return
|
||||
try:
|
||||
from transformers import AutoImageProcessor, AutoModelForObjectDetection
|
||||
|
||||
self._processor = AutoImageProcessor.from_pretrained(settings.detectors.tables.model)
|
||||
self._model = AutoModelForObjectDetection.from_pretrained(
|
||||
settings.detectors.tables.model
|
||||
)
|
||||
self._model.eval()
|
||||
self._initialized = True
|
||||
logger.info("table_detector_initialized", model=settings.detectors.tables.model)
|
||||
except Exception as e:
|
||||
logger.warning("table_detector_init_failed", error=str(e))
|
||||
self._initialized = True
|
||||
|
||||
def detect(self, image: np.ndarray) -> list[DetectedTable]:
|
||||
if not settings.detectors.tables.enabled:
|
||||
return []
|
||||
|
||||
self._init()
|
||||
if self._model is None or self._processor is None:
|
||||
return []
|
||||
|
||||
try:
|
||||
import torch
|
||||
|
||||
pil_image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
|
||||
pil_image = cv2.resize(pil_image, (800, 800))
|
||||
|
||||
inputs = self._processor(images=pil_image, return_tensors="pt")
|
||||
with torch.no_grad():
|
||||
outputs = self._model(**inputs)
|
||||
|
||||
target_sizes = torch.tensor([pil_image.shape[:2]])
|
||||
results = self._processor.post_process_object_detection(
|
||||
outputs,
|
||||
threshold=settings.detectors.tables.confidence_threshold,
|
||||
target_sizes=target_sizes,
|
||||
)[0]
|
||||
|
||||
tables = []
|
||||
for score, _label, box in zip(
|
||||
results["scores"], results["labels"], results["boxes"], strict=True
|
||||
):
|
||||
box = [float(x) for x in box]
|
||||
x1, y1, x2, y2 = map(int, box)
|
||||
crop = image[y1:y2, x1:x2]
|
||||
tables.append(
|
||||
DetectedTable(
|
||||
bbox=box,
|
||||
confidence=float(score),
|
||||
label="Table",
|
||||
image_crop=crop if crop.size > 0 else None,
|
||||
)
|
||||
)
|
||||
|
||||
logger.debug("tables_detected", count=len(tables))
|
||||
return tables
|
||||
except Exception as e:
|
||||
logger.warning("table_detection_failed", error=str(e))
|
||||
return []
|
||||
|
||||
|
||||
class TableStructureRecognizer:
|
||||
def __init__(self):
|
||||
self._model = None
|
||||
self._processor = None
|
||||
self._initialized = False
|
||||
|
||||
def _init(self):
|
||||
if self._initialized:
|
||||
return
|
||||
try:
|
||||
from transformers import AutoImageProcessor, AutoModelForObjectDetection
|
||||
|
||||
self._processor = AutoImageProcessor.from_pretrained(
|
||||
"microsoft/table-transformer-structure-recognition"
|
||||
)
|
||||
self._model = AutoModelForObjectDetection.from_pretrained(
|
||||
"microsoft/table-transformer-structure-recognition"
|
||||
)
|
||||
self._model.eval()
|
||||
self._initialized = True
|
||||
logger.info("table_structure_recognizer_initialized")
|
||||
except Exception as e:
|
||||
logger.warning("table_structure_recognizer_init_failed", error=str(e))
|
||||
self._initialized = True
|
||||
|
||||
def recognize(self, table_crop: np.ndarray) -> tuple[list[dict[str, Any]] | None, str | None]:
|
||||
if self._model is None or self._processor is None:
|
||||
return None, None
|
||||
|
||||
try:
|
||||
import torch
|
||||
|
||||
pil_image = cv2.cvtColor(table_crop, cv2.COLOR_BGR2RGB)
|
||||
inputs = self._processor(images=pil_image, return_tensors="pt")
|
||||
with torch.no_grad():
|
||||
outputs = self._model(**inputs)
|
||||
|
||||
target_sizes = torch.tensor([pil_image.shape[:2]])
|
||||
results = self._processor.post_process_object_detection(
|
||||
outputs, threshold=0.5, target_sizes=target_sizes
|
||||
)[0]
|
||||
|
||||
structure = []
|
||||
for score, label, box in zip(
|
||||
results["scores"], results["labels"], results["boxes"], strict=True
|
||||
):
|
||||
structure.append(
|
||||
{
|
||||
"label": self._model.config.id2label[int(label)],
|
||||
"confidence": float(score),
|
||||
"bbox": [float(x) for x in box],
|
||||
}
|
||||
)
|
||||
|
||||
markdown = self._structure_to_markdown(structure, table_crop.shape)
|
||||
return structure, markdown
|
||||
except Exception as e:
|
||||
logger.warning("table_structure_recognition_failed", error=str(e))
|
||||
return None, None
|
||||
|
||||
def _structure_to_markdown(self, structure: list[dict[str, Any]], shape: tuple) -> str | None:
|
||||
if not structure:
|
||||
return None
|
||||
|
||||
rows = {}
|
||||
for cell in structure:
|
||||
if cell["label"] in ("table row", "table cell"):
|
||||
y_center = (cell["bbox"][1] + cell["bbox"][3]) / 2
|
||||
x_center = (cell["bbox"][0] + cell["bbox"][2]) / 2
|
||||
row_idx = int(y_center / (shape[0] / 20))
|
||||
if row_idx not in rows:
|
||||
rows[row_idx] = []
|
||||
rows[row_idx].append((x_center, cell["label"]))
|
||||
|
||||
markdown_lines = []
|
||||
for row_idx in sorted(rows.keys()):
|
||||
cells = sorted(rows[row_idx], key=lambda x: x[0])
|
||||
row_cells = [c[1] for c in cells]
|
||||
markdown_lines.append("| " + " | ".join(row_cells) + " |")
|
||||
|
||||
if markdown_lines:
|
||||
header = markdown_lines[0]
|
||||
separator = "| " + " | ".join(["---"] * len(header.split("|")[1:-1])) + " |"
|
||||
markdown_lines.insert(1, separator)
|
||||
return "\n".join(markdown_lines)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def detect_tables(image: np.ndarray) -> list[DetectedTable]:
|
||||
detector = TableDetector()
|
||||
tables = detector.detect(image)
|
||||
|
||||
if not tables:
|
||||
return tables
|
||||
|
||||
recognizer = TableStructureRecognizer()
|
||||
for table in tables:
|
||||
if table.image_crop is not None:
|
||||
structure, markdown = recognizer.recognize(table.image_crop)
|
||||
table.structure = structure
|
||||
table.markdown = markdown
|
||||
|
||||
return tables
|
||||
Reference in New Issue
Block a user