improve the OCR pipeline processing and outputs formatting
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user