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:
188
src/ocr_pipeline/pipeline.py
Normal file
188
src/ocr_pipeline/pipeline.py
Normal file
@@ -0,0 +1,188 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from concurrent.futures import ProcessPoolExecutor, as_completed
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from ocr_pipeline.citations import extract_citations
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.detectors import detect_figures, detect_tables
|
||||
from ocr_pipeline.ocr import ImagePreprocessor, OCREngine, compute_image_hash
|
||||
from ocr_pipeline.ocr.parallel import ParallelProcessor
|
||||
from ocr_pipeline.output import MarkdownWriter
|
||||
from ocr_pipeline.postprocess import chunk_text, clean_text, extract_entities
|
||||
from ocr_pipeline.utils.db import get_db
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
from ocr_pipeline.watch import Watcher
|
||||
|
||||
try:
|
||||
from rich.console import Console
|
||||
|
||||
_console = Console()
|
||||
except ImportError:
|
||||
_console = None
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineResult:
|
||||
total_images: int = 0
|
||||
successful: int = 0
|
||||
failed: int = 0
|
||||
chunks_created: int = 0
|
||||
output_files: list[Path] = field(default_factory=list)
|
||||
errors: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
|
||||
class OCRPipeline:
|
||||
def __init__(self):
|
||||
self.engine = OCREngine()
|
||||
self.preprocessor = ImagePreprocessor()
|
||||
self.writer = MarkdownWriter()
|
||||
self.db = get_db()
|
||||
|
||||
def _process_single(self, image_path: Path) -> list[dict[str, Any]]:
|
||||
image = self.preprocessor.preprocess(image_path)
|
||||
result = self.engine.process(image)
|
||||
|
||||
image_hash = compute_image_hash(image_path)
|
||||
if self.db.is_processed(image_hash, image_path):
|
||||
logger.debug("already_processed", path=str(image_path))
|
||||
return []
|
||||
|
||||
cleaned = clean_text(result.text)
|
||||
chunks = chunk_text(cleaned.text)
|
||||
entities = extract_entities(cleaned.text)
|
||||
figures = detect_figures(image)
|
||||
tables = detect_tables(image)
|
||||
citations = extract_citations(cleaned.text)
|
||||
|
||||
output_files = []
|
||||
for i, chunk in enumerate(chunks):
|
||||
output_path = self.writer.write_chunk(
|
||||
source_path=image_path,
|
||||
source_hash=image_hash,
|
||||
text=chunk.content,
|
||||
entities=[e.text for e in entities.entities],
|
||||
figures=figures,
|
||||
tables=tables,
|
||||
citations=[f"{c.type}:{c.identifier}" for c in citations],
|
||||
chunk_index=i,
|
||||
total_chunks=len(chunks),
|
||||
ocr_engine=result.engine,
|
||||
ocr_confidence=result.confidence,
|
||||
language=result.language,
|
||||
)
|
||||
output_files.append(output_path)
|
||||
|
||||
self.db.mark_processed(
|
||||
file_hash=image_hash,
|
||||
path=str(image_path),
|
||||
timestamp=time.time(),
|
||||
engine=result.engine,
|
||||
confidence=result.confidence,
|
||||
chunks=len(chunks),
|
||||
run_id=0,
|
||||
)
|
||||
|
||||
return [{"path": p, "chunks": len(chunks)} for p in output_files]
|
||||
|
||||
def process_all(self, image_paths: list[Path] | None = None) -> PipelineResult:
|
||||
"""Process images through the full pipeline."""
|
||||
if image_paths is None:
|
||||
processor = ParallelProcessor()
|
||||
results = processor.process()
|
||||
image_paths = [r.path for r in results if r.success]
|
||||
else:
|
||||
# Validate paths exist
|
||||
image_paths = [p for p in image_paths if p.exists()]
|
||||
|
||||
result = PipelineResult(total_images=len(image_paths))
|
||||
|
||||
if not image_paths:
|
||||
return result
|
||||
|
||||
with ProcessPoolExecutor(max_workers=settings.processing.workers) as executor:
|
||||
future_to_path = {executor.submit(self._process_single, p): p for p in image_paths}
|
||||
|
||||
for future in as_completed(future_to_path):
|
||||
path = future_to_path[future]
|
||||
try:
|
||||
output_info = future.result(timeout=settings.processing.timeout_per_image)
|
||||
result.successful += 1
|
||||
result.chunks_created += sum(info["chunks"] for info in output_info)
|
||||
result.output_files.extend([info["path"] for info in output_info])
|
||||
except Exception as e:
|
||||
result.failed += 1
|
||||
result.errors.append({"path": str(path), "error": str(e)})
|
||||
logger.error("pipeline_failed", path=str(path), error=str(e))
|
||||
|
||||
return result
|
||||
|
||||
def watch(self):
|
||||
def process_callback(path: Path):
|
||||
self._process_single(path)
|
||||
|
||||
watcher = Watcher(process_callback)
|
||||
watcher.start()
|
||||
|
||||
try:
|
||||
while True:
|
||||
time.sleep(1)
|
||||
except KeyboardInterrupt:
|
||||
watcher.stop()
|
||||
|
||||
def get_stats(self) -> dict[str, Any]:
|
||||
db_stats = self.db.get_stats()
|
||||
return {
|
||||
"total_files_processed": db_stats.get("total_processed", 0),
|
||||
"total_runs": db_stats.get("total_runs", 0),
|
||||
"latest_run": db_stats.get("latest_run"),
|
||||
}
|
||||
|
||||
|
||||
def migrate_notebook(notebook_path: Path, pipeline: OCRPipeline) -> dict[str, Any]:
|
||||
import json
|
||||
|
||||
with open(notebook_path) as f:
|
||||
nb = json.load(f)
|
||||
|
||||
image_paths = set()
|
||||
for cell in nb.get("cells", []):
|
||||
if cell.get("cell_type") == "code":
|
||||
for output in cell.get("outputs", []):
|
||||
if output.get("output_type") == "stream":
|
||||
for line in output.get("text", []):
|
||||
parts = line.split()
|
||||
for p in parts:
|
||||
if "SCR-" in p and (
|
||||
p.endswith(".png") or p.endswith(".jpg") or p.endswith(".jpeg")
|
||||
):
|
||||
image_paths.add(p)
|
||||
|
||||
unique_paths = [Path(p) for p in image_paths if Path(p).exists()]
|
||||
|
||||
try:
|
||||
from rich.console import Console
|
||||
|
||||
console = Console()
|
||||
console.print(f"Found {len(unique_paths)} images in notebook")
|
||||
except ImportError:
|
||||
logger.info("Found %d images in notebook", len(unique_paths))
|
||||
|
||||
if not unique_paths:
|
||||
return {"total": 0, "successful": 0, "failed": 0, "errors": []}
|
||||
|
||||
result = pipeline.process_all(unique_paths)
|
||||
|
||||
return {
|
||||
"total": result.total_images,
|
||||
"successful": result.successful,
|
||||
"failed": result.failed,
|
||||
"chunks_created": result.chunks_created,
|
||||
"output_files": len(result.output_files),
|
||||
"errors": result.errors,
|
||||
}
|
||||
Reference in New Issue
Block a user