Files
kg-scr/src/ocr_pipeline/pipeline.py
Aman Nalakath a0211155fa refactor: solidify startup UX and engine-aware preprocessing
- Slim core deps: move ML stack to optional extras (paddle/tables/figures/scientific/full)
- Lazy settings proxy with config search paths (env var, cwd, user dir)
- New commands: init, demo, setup [basic|full], first-run guard on run/watch
- Engine-aware preprocessing: Paddle gets original image (fixes dark mode 0.83->0.95)
- Results table shows Skipped count; lazy run-dir creation
- kg_ocr marked experimental with extra, Docker defaults with OCR_PIPELINE_CONFIG
- 25/25 tests, ruff clean
2026-07-19 22:14:54 +02:00

165 lines
6.3 KiB
Python

from __future__ import annotations
import datetime
import time
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 Database, get_db
from ocr_pipeline.utils.logging import get_logger
from ocr_pipeline.watch import Watcher
logger = get_logger(__name__)
@dataclass
class PipelineResult:
total_images: int = 0
successful: int = 0
failed: int = 0
skipped: int = 0
chunks_created: int = 0
output_files: list[Path] = field(default_factory=list)
errors: list[dict[str, str]] = field(default_factory=list)
class OCRPipeline:
def __init__(self) -> None:
self.engine = OCREngine()
self.preprocessor = ImagePreprocessor()
self.writer = MarkdownWriter()
self.db: Database = get_db()
self._run_id: int | None = None
def _process_single(self, image_path: Path) -> list[Path]:
image_hash = compute_image_hash(image_path)
if settings.processing.skip_existing and self.db.is_processed(image_hash, image_path):
logger.info("already_processed", path=str(image_path))
return []
original_image = self.preprocessor.load_image(image_path)
# PaddleOCR reads the (size-capped) original; Tesseract the binarized derivative.
ocr_base = self.preprocessor.resize_if_needed(original_image)
preprocessed = self.preprocessor.preprocess_image(original_image)
ocr_result = self.engine.process(ocr_base, preprocessed)
cleaned = clean_text(ocr_result.text)
chunks = chunk_text(cleaned.text)
entities = extract_entities(cleaned.text)
# Detectors receive the original image: binarization/line removal damages their inputs.
figures = detect_figures(original_image)
tables = detect_tables(original_image)
citations = extract_citations(cleaned.text)
timestamp = datetime.datetime.now(datetime.UTC).isoformat()
output_files: list[Path] = []
for chunk in chunks:
metadata = {
"source_path": str(image_path),
"source_hash": image_hash,
"timestamp": timestamp,
"ocr_engine": ocr_result.engine,
"ocr_confidence_mean": ocr_result.confidence,
"language": ocr_result.language,
"entity_extraction_backend": entities.backend,
"detected_entities": [entity.text for entity in entities.entities],
"has_figures": bool(figures),
"has_tables": bool(tables),
"citations_found": [
f"{citation.type}:{citation.identifier}" for citation in citations
],
"chunk_index": chunk.chunk_index,
"total_chunks": len(chunks),
}
output_files.append(
self.writer.write_chunk(
content=chunk.content,
metadata=metadata,
figures=figures,
tables=tables,
entities=entities,
citations=citations,
)
)
self.db.mark_processed(
file_hash=image_hash,
path=str(image_path),
timestamp=time.time(),
engine=ocr_result.engine,
confidence=ocr_result.confidence,
chunks=len(chunks),
run_id=self._run_id,
)
return output_files
def process_single(self, image_path: Path) -> list[Path]:
"""Public single-image API used by notebook migration and watch mode."""
return self._process_single(image_path)
def process_all(self, image_paths: list[Path] | None = None) -> PipelineResult:
if image_paths is None:
image_paths = ParallelProcessor().find_images()
else:
image_paths = [path for path in image_paths if path.exists()]
result = PipelineResult(total_images=len(image_paths))
if not image_paths:
return result
run_timestamp = datetime.datetime.now(datetime.UTC).isoformat()
self._run_id, run_number = self.db.start_run()
self.writer.start_run(run_timestamp, run_number)
try:
# Output/database writes are deliberately centralized to avoid collisions and SQLite contention.
for path in image_paths:
try:
outputs = self._process_single(path)
if not outputs and settings.processing.skip_existing:
result.skipped += 1
else:
result.successful += 1
result.chunks_created += len(outputs)
result.output_files.extend(outputs)
except Exception as exc:
result.failed += 1
result.errors.append({"path": str(path), "error": str(exc)})
logger.exception("pipeline_failed", path=str(path))
if settings.output.write_consolidated:
consolidated = self.writer.write_consolidated()
if consolidated is not None:
result.output_files.append(consolidated)
finally:
self.db.complete_run(
self._run_id,
{
"total": result.total_images,
"successful": result.successful,
"failed": result.failed,
"chunks": result.chunks_created,
},
)
self._run_id = None
return result
def watch(self) -> None:
def process_callback(path: Path) -> None:
self.process_single(path)
watcher = Watcher(process_callback)
watcher.start()
try:
while True:
time.sleep(1)
except KeyboardInterrupt:
watcher.stop()
def get_stats(self, days: int = 30) -> dict[str, Any]:
return self.db.get_stats(days)