improve the OCR pipeline processing and outputs formatting

This commit is contained in:
2026-07-19 19:54:53 +02:00
parent d375db47c3
commit 1cf1c6eeab
30 changed files with 3838 additions and 3900 deletions

View File

@@ -1,7 +1,7 @@
from __future__ import annotations
import datetime
import time
from concurrent.futures import ProcessPoolExecutor, as_completed
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
@@ -13,17 +13,10 @@ 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.db import Database, 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__)
@@ -32,157 +25,104 @@ 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, Any]] = field(default_factory=list)
errors: list[dict[str, str]] = field(default_factory=list)
class OCRPipeline:
def __init__(self):
def __init__(self) -> None:
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)
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 self.db.is_processed(image_hash, image_path):
logger.debug("already_processed", path=str(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 []
cleaned = clean_text(result.text)
original_image = self.preprocessor.load_image(image_path)
ocr_image = self.preprocessor.preprocess_image(original_image)
ocr_result = self.engine.process(ocr_image)
cleaned = clean_text(ocr_result.text)
chunks = chunk_text(cleaned.text)
entities = extract_entities(cleaned.text)
figures = detect_figures(image)
tables = detect_tables(image)
# 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 = []
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)
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=result.engine,
confidence=result.confidence,
chunks=len(chunks),
run_id=0,
)
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
return [{"path": p, "chunks": len(chunks)} for p in 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:
"""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]
image_paths = ParallelProcessor().find_images()
else:
# Validate paths exist
image_paths = [p for p in image_paths if p.exists()]
image_paths = [path for path in image_paths if path.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]
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:
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:
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(e)})
logger.error("pipeline_failed", path=str(path), error=str(e))
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):
def process_callback(path: Path):
self._process_single(path)
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) -> 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,
}
def get_stats(self, days: int = 30) -> dict[str, Any]:
return self.db.get_stats(days)