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:
26
.env.example
Normal file
26
.env.example
Normal file
@@ -0,0 +1,26 @@
|
||||
# OCR Pipeline Configuration
|
||||
# Copy to .env and modify as needed
|
||||
|
||||
# Logging
|
||||
LOG_LEVEL=INFO
|
||||
LOG_FILE=./data/logs/ocr_pipeline.log
|
||||
|
||||
# Input paths (override config.yaml)
|
||||
INPUT_PATHS=~/Pictures,/mnt/storage3/aman/screenshots
|
||||
|
||||
# Processing
|
||||
WORKERS=4
|
||||
OCR_ENGINE=paddleocr
|
||||
USE_GPU=false
|
||||
|
||||
# Output
|
||||
OUTPUT_DIR=./data/ocr_output
|
||||
ORGANIZE_BY=date_run
|
||||
|
||||
# Watch mode
|
||||
WATCH_ENABLED=true
|
||||
WATCH_DB_PATH=./data/processed_files.db
|
||||
|
||||
# Optional: GROBID for citation parsing
|
||||
GROBID_ENABLED=false
|
||||
GROBID_URL=http://localhost:8070
|
||||
27
.pre-commit-config.yaml
Normal file
27
.pre-commit-config.yaml
Normal file
@@ -0,0 +1,27 @@
|
||||
repos:
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.2.0
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--fix]
|
||||
- id: ruff-format
|
||||
|
||||
- repo: https://github.com/pre-commit/mirrors-mypy
|
||||
rev: v1.8.0
|
||||
hooks:
|
||||
- id: mypy
|
||||
additional_dependencies: [pydantic, pyyaml, typer, rich, structlog, numpy, opencv-python-headless]
|
||||
|
||||
- repo: https://github.com/pre-commit/pre-commit-hooks
|
||||
rev: v4.5.0
|
||||
hooks:
|
||||
- id: trailing-whitespace
|
||||
- id: end-of-file-fixer
|
||||
- id: check-yaml
|
||||
- id: check-toml
|
||||
- id: check-added-large-files
|
||||
|
||||
- repo: https://github.com/pycqa/isort
|
||||
rev: 5.13.2
|
||||
hooks:
|
||||
- id: isort
|
||||
26
Dockerfile
Normal file
26
Dockerfile
Normal file
@@ -0,0 +1,26 @@
|
||||
FROM python:3.11-slim
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
tesseract-ocr \
|
||||
libtesseract-dev \
|
||||
poppler-utils \
|
||||
libgl1 \
|
||||
libglib2.0-0 \
|
||||
libsm6 \
|
||||
libxext6 \
|
||||
libxrender-dev \
|
||||
libgomp1 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY pyproject.toml ./
|
||||
RUN pip install --no-cache-dir uv && uv sync --frozen
|
||||
|
||||
COPY src ./src
|
||||
COPY config.yaml ./
|
||||
|
||||
ENV PYTHONPATH=/app/src
|
||||
ENV OMP_NUM_THREADS=4
|
||||
|
||||
ENTRYPOINT ["ocr-pipeline"]
|
||||
175
README.md
175
README.md
@@ -0,0 +1,175 @@
|
||||
# OCR Pipeline for Life Science Screenshots
|
||||
|
||||
Turns scientific screenshots into RAG-ready Markdown, with metadata attached.
|
||||
|
||||
- OCR via PaddleOCR, falls back to Tesseract
|
||||
- Preprocessing: deskew, denoise, CLAHE, line removal
|
||||
- Figure/table detection
|
||||
- Entity extraction with scispaCy
|
||||
- Citation matching
|
||||
- Chunking for embedding
|
||||
|
||||
## Quick start
|
||||
|
||||
```bash
|
||||
uv sync # or: pip install -e .
|
||||
pre-commit install
|
||||
|
||||
ocr-pipeline run # process screenshots
|
||||
ocr-pipeline watch # watch a folder
|
||||
ocr-pipeline reprocess ocr_sc.ipynb # pull screenshots out of a notebook
|
||||
ocr-pipeline status
|
||||
```
|
||||
|
||||
## Docker
|
||||
|
||||
```bash
|
||||
docker build -t ocr-pipeline .
|
||||
|
||||
docker run -v ~/Pictures:/data/screenshots -v ./data:/app/data ocr-pipeline run
|
||||
docker run -v ~/Pictures:/data/screenshots -v ./data:/app/data ocr-pipeline watch
|
||||
```
|
||||
|
||||
## Config
|
||||
|
||||
<details>
|
||||
<summary><code>config.yaml</code></summary>
|
||||
|
||||
```yaml
|
||||
input:
|
||||
paths: ["~/Pictures", "/mnt/storage3/aman/screenshots"]
|
||||
patterns: ["SCR-*.png", "*.jpg", "*.jpeg", "*.tiff"]
|
||||
recursive: true
|
||||
|
||||
ocr:
|
||||
engine: "paddleocr" # paddleocr | tesseract | auto
|
||||
languages: ["en", "latin"]
|
||||
use_gpu: false
|
||||
preprocess:
|
||||
deskew: true
|
||||
denoise: true
|
||||
clahe: true
|
||||
adaptive_threshold: true
|
||||
remove_lines: true
|
||||
|
||||
processing:
|
||||
workers: 4
|
||||
batch_size: 10
|
||||
retry_attempts: 2
|
||||
|
||||
detectors:
|
||||
figures:
|
||||
enabled: true
|
||||
confidence_threshold: 0.7
|
||||
tables:
|
||||
enabled: true
|
||||
confidence_threshold: 0.7
|
||||
|
||||
entities:
|
||||
enabled: true
|
||||
model: "en_core_sci_lg"
|
||||
|
||||
citations:
|
||||
regex_enabled: true
|
||||
grobid_enabled: false # set true if you're running a GROBID server
|
||||
|
||||
chunking:
|
||||
chunk_size: 1000
|
||||
chunk_overlap: 200
|
||||
|
||||
output:
|
||||
base_directory: "./data/ocr_output"
|
||||
organize_by: "date_run" # date_run | source_dir | flat
|
||||
frontmatter:
|
||||
- source_path
|
||||
- source_hash
|
||||
- timestamp
|
||||
- ocr_engine
|
||||
- ocr_confidence_mean
|
||||
- language
|
||||
- detected_entities
|
||||
- has_figures
|
||||
- has_tables
|
||||
- citations_found
|
||||
- chunk_index
|
||||
- total_chunks
|
||||
|
||||
watch:
|
||||
enabled: true
|
||||
debounce_seconds: 5
|
||||
db_path: "./data/processed_files.db"
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
## Output
|
||||
|
||||
Each chunk is a Markdown file with YAML frontmatter:
|
||||
|
||||
```markdown
|
||||
---
|
||||
source_path: "/Users/Aman/Pictures/SCR-20250115-gel.png"
|
||||
source_hash: "a1b2c3d4e5f6..."
|
||||
timestamp: "2025-01-15T10:30:00Z"
|
||||
ocr_engine: "paddleocr"
|
||||
ocr_confidence_mean: 0.91
|
||||
language: "en"
|
||||
detected_entities: ["GENE", "PROTEIN", "CHEMICAL"]
|
||||
has_figures: true
|
||||
has_tables: false
|
||||
citations_found: ["DOI:10.1038/nature12345", "PMID:12345678"]
|
||||
chunk_index: 0
|
||||
total_chunks: 2
|
||||
---
|
||||
|
||||
# Screenshot: SCR-20250115-gel.png
|
||||
|
||||
## Figures
|
||||
|
||||
### Figure 1
|
||||
- BBox: [100, 200, 800, 600]
|
||||
- Confidence: 0.92
|
||||
- Caption: "Western blot showing BRCA1 expression..."
|
||||
|
||||
## Detected Entities
|
||||
|
||||
- BRCA1
|
||||
- CRISPR
|
||||
- β-actin
|
||||
|
||||
## Citations
|
||||
|
||||
- DOI: 10.1038/nature12345
|
||||
- PMID: 12345678
|
||||
|
||||
## Extracted Text
|
||||
|
||||
**Western Blot Analysis of BRCA1 Expression**
|
||||
|
||||
Lane 1: WT control
|
||||
Lane 2: BRCA1 KO (CRISPR)
|
||||
Lane 3: BRCA1 KO + pBRCA1-WT rescue
|
||||
Lane 4: BRCA1 KO + pBRCA1-C61G mutant
|
||||
|
||||
Anti-BRCA1 (1:1000), Anti-β-actin (1:5000)
|
||||
```
|
||||
|
||||
Files land in `data/ocr_output/<date>/run_NNN/`.
|
||||
|
||||
## Life science specifics
|
||||
|
||||
Gene/protein names go through scispaCy's `en_core_sci_lg`. Chemical formulas and units (µM, ng/mL, kb/Mb/Gb, °C, ×g) get normalized, scientific notation gets cleaned up (`1.5×10⁻³` → `1.5×10^-3`), and gel/blot figures get their captions pulled out separately. Citations are matched by regex for DOI, PMID, arXiv, PMC, and ISBN.
|
||||
|
||||
## Models
|
||||
|
||||
First run downloads what it needs, cached in `~/.cache/ocr_pipeline/`:
|
||||
|
||||
- PaddleOCR models (~200MB)
|
||||
- scispaCy `en_core_sci_lg` (~800MB)
|
||||
- Table Transformer (~500MB)
|
||||
- LayoutParser PubLayNet (~300MB)
|
||||
|
||||
## TODO
|
||||
|
||||
- [ ] Build a knowledge graph from extracted entities/citations
|
||||
|
||||
|
||||
105
config.yaml
Normal file
105
config.yaml
Normal file
@@ -0,0 +1,105 @@
|
||||
input:
|
||||
paths:
|
||||
- "~/Pictures"
|
||||
- "/mnt/storage3/aman/screenshots"
|
||||
patterns:
|
||||
- "SCR-*.png"
|
||||
- "*.jpg"
|
||||
- "*.jpeg"
|
||||
- "*.tiff"
|
||||
- "*.bmp"
|
||||
recursive: true
|
||||
|
||||
ocr:
|
||||
engine: "paddleocr" # "paddleocr" | "tesseract" | "auto"
|
||||
languages: ["en"] # For tesseract: eng, lat, deu, fra, spa, ita, por, chi_sim, jpn, kor
|
||||
use_gpu: false
|
||||
use_angle_cls: true
|
||||
det_db_thresh: 0.3
|
||||
det_db_box_thresh: 0.6
|
||||
det_db_unclip_ratio: 1.5
|
||||
rec_batch_num: 6
|
||||
cpu_threads: 4
|
||||
preprocess:
|
||||
deskew: true
|
||||
denoise: true
|
||||
clahe: true
|
||||
adaptive_threshold: true
|
||||
remove_lines: true
|
||||
max_dimension: 4096
|
||||
|
||||
processing:
|
||||
workers: 4
|
||||
batch_size: 10
|
||||
retry_attempts: 2
|
||||
timeout_per_image: 120
|
||||
skip_existing: true
|
||||
|
||||
detectors:
|
||||
figures:
|
||||
enabled: true
|
||||
model: "lp://PubLayNet/faster_rcnn_R_50_FPN_3x/config"
|
||||
confidence_threshold: 0.7
|
||||
tables:
|
||||
enabled: true
|
||||
model: "microsoft/table-transformer-detection"
|
||||
confidence_threshold: 0.7
|
||||
captions:
|
||||
enabled: true
|
||||
model: "lp://PubLayNet/faster_rcnn_R_50_FPN_3x/config"
|
||||
confidence_threshold: 0.5
|
||||
|
||||
entities:
|
||||
enabled: true
|
||||
model: "en_core_sci_lg"
|
||||
types:
|
||||
- "GENE"
|
||||
- "PROTEIN"
|
||||
- "CHEMICAL"
|
||||
- "SPECIES"
|
||||
- "DISEASE"
|
||||
- "CELL_LINE"
|
||||
- "ORGANISM"
|
||||
- "CELL_TYPE"
|
||||
merge_entities: true
|
||||
|
||||
citations:
|
||||
regex_enabled: true
|
||||
grobid_enabled: false
|
||||
grobid_url: "http://localhost:8070"
|
||||
grobid_timeout: 30
|
||||
|
||||
chunking:
|
||||
chunk_size: 1000
|
||||
chunk_overlap: 200
|
||||
separators: ["\n\n", "\n", ". ", " ", ""]
|
||||
keep_separator: true
|
||||
|
||||
output:
|
||||
base_directory: "./data/ocr_output"
|
||||
format: "markdown"
|
||||
organize_by: "date_run" # "date_run" | "source_dir" | "flat"
|
||||
frontmatter:
|
||||
- "source_path"
|
||||
- "source_hash"
|
||||
- "timestamp"
|
||||
- "ocr_engine"
|
||||
- "ocr_confidence_mean"
|
||||
- "language"
|
||||
- "detected_entities"
|
||||
- "has_figures"
|
||||
- "has_tables"
|
||||
- "citations_found"
|
||||
- "chunk_index"
|
||||
- "total_chunks"
|
||||
|
||||
watch:
|
||||
enabled: true
|
||||
debounce_seconds: 5
|
||||
ignore_patterns:
|
||||
- ".DS_Store"
|
||||
- "*.tmp"
|
||||
- "*.partial"
|
||||
- "*.crdownload"
|
||||
db_path: "./data/processed_files.db"
|
||||
poll_interval: 1.0
|
||||
134
pyproject.toml
Normal file
134
pyproject.toml
Normal file
@@ -0,0 +1,134 @@
|
||||
[project]
|
||||
name = "ocr-pipeline"
|
||||
version = "0.1.0"
|
||||
description = "OCR pipeline for life science screenshots - RAG ready"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
# Core
|
||||
"typer[all]>=0.9.0",
|
||||
"pydantic>=2.6.0",
|
||||
"pydantic-settings>=2.2.0",
|
||||
"pyyaml>=6.0.1",
|
||||
"structlog>=24.1.0",
|
||||
"rich>=13.7.0",
|
||||
"tqdm>=4.66.0",
|
||||
|
||||
# Image processing
|
||||
"opencv-python-headless>=4.9.0",
|
||||
"pillow>=10.2.0",
|
||||
"numpy>=1.26.0",
|
||||
"scikit-image>=0.22.0",
|
||||
|
||||
# OCR engines
|
||||
"paddleocr>=2.7.0",
|
||||
"pytesseract>=0.3.10",
|
||||
|
||||
# Table detection
|
||||
"torch>=2.2.0",
|
||||
"torchvision>=0.17.0",
|
||||
"transformers>=4.38.0",
|
||||
"layoutparser>=0.3.0",
|
||||
|
||||
# NLP
|
||||
"langchain-text-splitters>=0.0.2",
|
||||
"spacy>=3.7.0",
|
||||
|
||||
# Database
|
||||
"sqlite-utils>=3.37.0",
|
||||
|
||||
# File watching
|
||||
"watchdog>=3.0.0",
|
||||
|
||||
# Utilities
|
||||
"python-slugify>=8.0.0",
|
||||
"xxhash>=3.4.0",
|
||||
"python-magic>=0.4.27",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = [
|
||||
"pytest>=8.0.0",
|
||||
"pytest-cov>=4.1.0",
|
||||
"pytest-mock>=3.12.0",
|
||||
"ruff>=0.2.0",
|
||||
"mypy>=1.8.0",
|
||||
"pre-commit>=3.6.0",
|
||||
]
|
||||
grobid = [
|
||||
"grobid-client>=0.8.0",
|
||||
]
|
||||
full = [
|
||||
"ocr-pipeline[dev]",
|
||||
"ocr-pipeline[grobid]",
|
||||
]
|
||||
|
||||
[project.entry-points.console_scripts]
|
||||
ocr-pipeline = "ocr_pipeline.cli:app"
|
||||
|
||||
[tool.uv]
|
||||
dev-dependencies = [
|
||||
"pytest>=8.0.0",
|
||||
"pytest-cov>=4.1.0",
|
||||
"pytest-mock>=3.12.0",
|
||||
"ruff>=0.2.0",
|
||||
"mypy>=1.8.0",
|
||||
"pre-commit>=3.6.0",
|
||||
]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 100
|
||||
target-version = "py311"
|
||||
select = [
|
||||
"E", # pycodestyle
|
||||
"F", # pyflakes
|
||||
"I", # isort
|
||||
"N", # pep8-naming
|
||||
"W", # pycodestyle warnings
|
||||
"UP", # pyupgrade
|
||||
"B", # flake8-bugbear
|
||||
"C4", # flake8-comprehensions
|
||||
"T20", # flake8-print
|
||||
]
|
||||
ignore = [
|
||||
"E501", # line too long (handled by formatter)
|
||||
"B008", # function calls in default argument
|
||||
]
|
||||
fixable = ["I", "E", "W", "F"]
|
||||
unfixable = ["N", "B", "C4", "T20", "UP"]
|
||||
|
||||
[tool.ruff.format]
|
||||
quote-style = "double"
|
||||
indent-style = "space"
|
||||
skip-magic-trailing-comma = false
|
||||
line-ending = "lf"
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.11"
|
||||
warn_return_any = true
|
||||
warn_unused_configs = true
|
||||
disallow_untyped_defs = true
|
||||
disallow_incomplete_defs = true
|
||||
check_untyped_defs = true
|
||||
no_implicit_optional = true
|
||||
strict_equality = true
|
||||
show_error_codes = true
|
||||
pretty = true
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
python_files = ["test_*.py"]
|
||||
python_classes = ["Test*"]
|
||||
python_functions = ["test_*"]
|
||||
addopts = "-v --tb=short"
|
||||
filterwarnings = [
|
||||
"ignore::DeprecationWarning",
|
||||
"ignore::UserWarning:torch",
|
||||
"ignore::UserWarning:transformers",
|
||||
]
|
||||
|
||||
[tool.pre-commit-hooks]
|
||||
ruff = { id = "ruff", args = ["check", "--fix"] }
|
||||
ruff-format = { id = "ruff", args = ["format"] }
|
||||
mypy = { id = "mypy" }
|
||||
pytest = { id = "pytest", args = ["-x"] }
|
||||
276
src/ocr_pipeline.egg-info/PKG-INFO
Normal file
276
src/ocr_pipeline.egg-info/PKG-INFO
Normal file
@@ -0,0 +1,276 @@
|
||||
Metadata-Version: 2.4
|
||||
Name: ocr-pipeline
|
||||
Version: 0.1.0
|
||||
Summary: OCR pipeline for life science screenshots - RAG ready
|
||||
Requires-Python: >=3.11
|
||||
Description-Content-Type: text/markdown
|
||||
Requires-Dist: typer[all]>=0.9.0
|
||||
Requires-Dist: pydantic>=2.6.0
|
||||
Requires-Dist: pydantic-settings>=2.2.0
|
||||
Requires-Dist: pyyaml>=6.0.1
|
||||
Requires-Dist: structlog>=24.1.0
|
||||
Requires-Dist: rich>=13.7.0
|
||||
Requires-Dist: tqdm>=4.66.0
|
||||
Requires-Dist: opencv-python-headless>=4.9.0
|
||||
Requires-Dist: pillow>=10.2.0
|
||||
Requires-Dist: numpy>=1.26.0
|
||||
Requires-Dist: scikit-image>=0.22.0
|
||||
Requires-Dist: paddleocr>=2.7.0
|
||||
Requires-Dist: pytesseract>=0.3.10
|
||||
Requires-Dist: torch>=2.2.0
|
||||
Requires-Dist: torchvision>=0.17.0
|
||||
Requires-Dist: transformers>=4.38.0
|
||||
Requires-Dist: layoutparser>=0.3.0
|
||||
Requires-Dist: langchain-text-splitters>=0.0.2
|
||||
Requires-Dist: spacy>=3.7.0
|
||||
Requires-Dist: sqlite-utils>=3.37.0
|
||||
Requires-Dist: watchdog>=3.0.0
|
||||
Requires-Dist: python-slugify>=8.0.0
|
||||
Requires-Dist: xxhash>=3.4.0
|
||||
Requires-Dist: python-magic>=0.4.27
|
||||
Provides-Extra: dev
|
||||
Requires-Dist: pytest>=8.0.0; extra == "dev"
|
||||
Requires-Dist: pytest-cov>=4.1.0; extra == "dev"
|
||||
Requires-Dist: pytest-mock>=3.12.0; extra == "dev"
|
||||
Requires-Dist: ruff>=0.2.0; extra == "dev"
|
||||
Requires-Dist: mypy>=1.8.0; extra == "dev"
|
||||
Requires-Dist: pre-commit>=3.6.0; extra == "dev"
|
||||
Provides-Extra: grobid
|
||||
Requires-Dist: grobid-client>=0.8.0; extra == "grobid"
|
||||
Provides-Extra: full
|
||||
Requires-Dist: ocr-pipeline[dev]; extra == "full"
|
||||
Requires-Dist: ocr-pipeline[grobid]; extra == "full"
|
||||
|
||||
# OCR Pipeline for Life Science Screenshots
|
||||
|
||||
A production-ready OCR pipeline that extracts text from scientific screenshots and converts them into RAG-ready Markdown files with rich metadata.
|
||||
|
||||
## Features
|
||||
|
||||
- **Multi-engine OCR**: PaddleOCR (primary) with Tesseract fallback
|
||||
- **Image preprocessing**: Deskewing, denoising, CLAHE contrast enhancement, line removal
|
||||
- **Parallel processing**: Multi-process worker pool for throughput
|
||||
- **Scientific entity recognition**: Genes, proteins, chemicals, species, diseases, cell lines via scispaCy
|
||||
- **Figure & table detection**: LayoutParser + Table Transformer
|
||||
- **Citation extraction**: DOI, PMID, arXiv, PMC, ISBN via regex + optional GROBID
|
||||
- **Semantic chunking**: LangChain recursive splitter optimized for scientific text
|
||||
- **RAG-ready output**: Markdown with YAML frontmatter (source, hash, timestamp, entities, confidence)
|
||||
- **Watch mode**: File system monitoring with SQLite persistence for incremental processing
|
||||
- **Notebook migration**: Reprocess screenshots from existing Jupyter notebooks
|
||||
- **Life-science focused**: Handles scientific notation, units, gene symbols, chemical formulas
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
# Install with uv (recommended)
|
||||
uv sync
|
||||
|
||||
# Or with pip
|
||||
pip install -e .
|
||||
|
||||
# Install pre-commit hooks
|
||||
pre-commit install
|
||||
|
||||
# Run on your screenshots
|
||||
ocr-pipeline run
|
||||
|
||||
# Watch for new screenshots
|
||||
ocr-pipeline watch
|
||||
|
||||
# Reprocess from notebook
|
||||
ocr-pipeline reprocess ocr_sc.ipynb
|
||||
|
||||
# Check status
|
||||
ocr-pipeline status
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
Edit `config.yaml`:
|
||||
|
||||
```yaml
|
||||
input:
|
||||
paths: ["~/Pictures", "/mnt/storage3/aman/screenshots"]
|
||||
patterns: ["SCR-*.png", "*.jpg", "*.jpeg", "*.tiff"]
|
||||
recursive: true
|
||||
|
||||
ocr:
|
||||
engine: "paddleocr" # paddleocr | tesseract | auto
|
||||
languages: ["en", "latin"]
|
||||
use_gpu: false
|
||||
preprocess:
|
||||
deskew: true
|
||||
denoise: true
|
||||
clahe: true
|
||||
adaptive_threshold: true
|
||||
remove_lines: true
|
||||
|
||||
processing:
|
||||
workers: 4
|
||||
batch_size: 10
|
||||
retry_attempts: 2
|
||||
|
||||
detectors:
|
||||
figures:
|
||||
enabled: true
|
||||
confidence_threshold: 0.7
|
||||
tables:
|
||||
enabled: true
|
||||
confidence_threshold: 0.7
|
||||
|
||||
entities:
|
||||
enabled: true
|
||||
model: "en_core_sci_lg" # scispaCy large model
|
||||
|
||||
citations:
|
||||
regex_enabled: true
|
||||
grobid_enabled: false # Set true if running GROBID server
|
||||
|
||||
chunking:
|
||||
chunk_size: 1000
|
||||
chunk_overlap: 200
|
||||
|
||||
output:
|
||||
base_directory: "./data/ocr_output"
|
||||
organize_by: "date_run" # date_run | source_dir | flat
|
||||
frontmatter:
|
||||
- source_path
|
||||
- source_hash
|
||||
- timestamp
|
||||
- ocr_engine
|
||||
- ocr_confidence_mean
|
||||
- language
|
||||
- detected_entities
|
||||
- has_figures
|
||||
- has_tables
|
||||
- citations_found
|
||||
- chunk_index
|
||||
- total_chunks
|
||||
|
||||
watch:
|
||||
enabled: true
|
||||
debounce_seconds: 5
|
||||
db_path: "./data/processed_files.db"
|
||||
```
|
||||
|
||||
## Output Format
|
||||
|
||||
Each chunk produces a Markdown file with YAML frontmatter:
|
||||
|
||||
```markdown
|
||||
---
|
||||
source_path: "/Users/Aman/Pictures/SCR-20250115-gel.png"
|
||||
source_hash: "a1b2c3d4e5f6..."
|
||||
timestamp: "2025-01-15T10:30:00Z"
|
||||
ocr_engine: "paddleocr"
|
||||
ocr_confidence_mean: 0.91
|
||||
language: "en"
|
||||
detected_entities: ["GENE", "PROTEIN", "CHEMICAL"]
|
||||
has_figures: true
|
||||
has_tables: false
|
||||
citations_found: ["DOI:10.1038/nature12345", "PMID:12345678"]
|
||||
chunk_index: 0
|
||||
total_chunks: 2
|
||||
---
|
||||
|
||||
# Screenshot: SCR-20250115-gel.png
|
||||
|
||||
## Figures
|
||||
|
||||
### Figure 1
|
||||
- **BBox**: [100, 200, 800, 600]
|
||||
- **Confidence**: 0.92
|
||||
- **Caption**: "Western blot showing BRCA1 expression..."
|
||||
|
||||
## Detected Entities
|
||||
|
||||
- BRCA1
|
||||
- CRISPR
|
||||
- β-actin
|
||||
|
||||
## Citations
|
||||
|
||||
- DOI: 10.1038/nature12345
|
||||
- PMID: 12345678
|
||||
|
||||
## Extracted Text
|
||||
|
||||
**Western Blot Analysis of BRCA1 Expression**
|
||||
|
||||
Lane 1: WT control
|
||||
Lane 2: BRCA1 KO (CRISPR)
|
||||
Lane 3: BRCA1 KO + pBRCA1-WT rescue
|
||||
Lane 4: BRCA1 KO + pBRCA1-C61G mutant
|
||||
|
||||
Anti-BRCA1 (1:1000), Anti-β-actin (1:5000)
|
||||
```
|
||||
|
||||
## Directory Structure
|
||||
|
||||
```
|
||||
data/ocr_output/
|
||||
├── 2025-01-15/
|
||||
│ ├── run_001/
|
||||
│ │ ├── SCR-20250115-gel_chunk_000.md
|
||||
│ │ └── SCR-20250115-gel_chunk_001.md
|
||||
│ └── run_002/
|
||||
│ └── ...
|
||||
└── 2025-01-16/
|
||||
└── run_001/
|
||||
└── ...
|
||||
```
|
||||
|
||||
## Docker
|
||||
|
||||
```bash
|
||||
# Build
|
||||
docker build -t ocr-pipeline .
|
||||
|
||||
# Run once
|
||||
docker run -v ~/Pictures:/data/screenshots -v ./data:/app/data ocr-pipeline run
|
||||
|
||||
# Watch mode
|
||||
docker run -v ~/Pictures:/data/screenshots -v ./data:/app/data ocr-pipeline watch
|
||||
```
|
||||
|
||||
## Life Science Optimizations
|
||||
|
||||
| Feature | Implementation |
|
||||
|---------|----------------|
|
||||
| Gene/Protein names | scispaCy `en_core_sci_lg` NER |
|
||||
| Chemical formulas | Regex + unit normalization (µM, ng/mL, kb, etc.) |
|
||||
| Scientific notation | `1.5×10⁻³` → `1.5×10^-3` |
|
||||
| Gel/blot lanes | Figure detection + caption extraction |
|
||||
| Citations | DOI, PMID, arXiv, PMC, ISBN patterns |
|
||||
| Units | µM, ng/mL, kb/Mb/Gb, °C, ×g, etc. |
|
||||
|
||||
## Development
|
||||
|
||||
```bash
|
||||
# Install dev dependencies
|
||||
uv sync --dev
|
||||
|
||||
# Run tests
|
||||
pytest
|
||||
|
||||
# Lint
|
||||
ruff check .
|
||||
ruff format .
|
||||
|
||||
# Type check
|
||||
mypy src/
|
||||
```
|
||||
|
||||
## Model Downloads
|
||||
|
||||
First run downloads models automatically:
|
||||
- PaddleOCR detection/recognition models (~200MB)
|
||||
- scispaCy `en_core_sci_lg` (~800MB)
|
||||
- Table Transformer (~500MB)
|
||||
- LayoutParser PubLayNet (~300MB)
|
||||
|
||||
Cache location: `~/.cache/ocr_pipeline/`
|
||||
|
||||
## License
|
||||
|
||||
MIT
|
||||
35
src/ocr_pipeline.egg-info/SOURCES.txt
Normal file
35
src/ocr_pipeline.egg-info/SOURCES.txt
Normal file
@@ -0,0 +1,35 @@
|
||||
README.md
|
||||
pyproject.toml
|
||||
src/ocr_pipeline/__init__.py
|
||||
src/ocr_pipeline/cli.py
|
||||
src/ocr_pipeline/config.py
|
||||
src/ocr_pipeline/pipeline.py
|
||||
src/ocr_pipeline/watch.py
|
||||
src/ocr_pipeline.egg-info/PKG-INFO
|
||||
src/ocr_pipeline.egg-info/SOURCES.txt
|
||||
src/ocr_pipeline.egg-info/dependency_links.txt
|
||||
src/ocr_pipeline.egg-info/entry_points.txt
|
||||
src/ocr_pipeline.egg-info/requires.txt
|
||||
src/ocr_pipeline.egg-info/top_level.txt
|
||||
src/ocr_pipeline/citations/__init__.py
|
||||
src/ocr_pipeline/citations/extractor.py
|
||||
src/ocr_pipeline/citations/regex_extractor.py
|
||||
src/ocr_pipeline/detectors/__init__.py
|
||||
src/ocr_pipeline/detectors/figures.py
|
||||
src/ocr_pipeline/detectors/tables.py
|
||||
src/ocr_pipeline/ocr/__init__.py
|
||||
src/ocr_pipeline/ocr/engine.py
|
||||
src/ocr_pipeline/ocr/parallel.py
|
||||
src/ocr_pipeline/output/__init__.py
|
||||
src/ocr_pipeline/output/markdown.py
|
||||
src/ocr_pipeline/postprocess/__init__.py
|
||||
src/ocr_pipeline/postprocess/chunk.py
|
||||
src/ocr_pipeline/postprocess/clean.py
|
||||
src/ocr_pipeline/postprocess/entities.py
|
||||
src/ocr_pipeline/utils/__init__.py
|
||||
src/ocr_pipeline/utils/db.py
|
||||
src/ocr_pipeline/utils/logging.py
|
||||
src/ocr_pipeline/utils/migrate.py
|
||||
src/ocr_pipeline/watch/__init__.py
|
||||
src/ocr_pipeline/watch/watcher.py
|
||||
tests/test_postprocess.py
|
||||
1
src/ocr_pipeline.egg-info/dependency_links.txt
Normal file
1
src/ocr_pipeline.egg-info/dependency_links.txt
Normal file
@@ -0,0 +1 @@
|
||||
|
||||
2
src/ocr_pipeline.egg-info/entry_points.txt
Normal file
2
src/ocr_pipeline.egg-info/entry_points.txt
Normal file
@@ -0,0 +1,2 @@
|
||||
[console_scripts]
|
||||
ocr-pipeline = ocr_pipeline.cli:app
|
||||
39
src/ocr_pipeline.egg-info/requires.txt
Normal file
39
src/ocr_pipeline.egg-info/requires.txt
Normal file
@@ -0,0 +1,39 @@
|
||||
typer[all]>=0.9.0
|
||||
pydantic>=2.6.0
|
||||
pydantic-settings>=2.2.0
|
||||
pyyaml>=6.0.1
|
||||
structlog>=24.1.0
|
||||
rich>=13.7.0
|
||||
tqdm>=4.66.0
|
||||
opencv-python-headless>=4.9.0
|
||||
pillow>=10.2.0
|
||||
numpy>=1.26.0
|
||||
scikit-image>=0.22.0
|
||||
paddleocr>=2.7.0
|
||||
pytesseract>=0.3.10
|
||||
torch>=2.2.0
|
||||
torchvision>=0.17.0
|
||||
transformers>=4.38.0
|
||||
layoutparser>=0.3.0
|
||||
langchain-text-splitters>=0.0.2
|
||||
spacy>=3.7.0
|
||||
sqlite-utils>=3.37.0
|
||||
watchdog>=3.0.0
|
||||
python-slugify>=8.0.0
|
||||
xxhash>=3.4.0
|
||||
python-magic>=0.4.27
|
||||
|
||||
[dev]
|
||||
pytest>=8.0.0
|
||||
pytest-cov>=4.1.0
|
||||
pytest-mock>=3.12.0
|
||||
ruff>=0.2.0
|
||||
mypy>=1.8.0
|
||||
pre-commit>=3.6.0
|
||||
|
||||
[full]
|
||||
ocr-pipeline[dev]
|
||||
ocr-pipeline[grobid]
|
||||
|
||||
[grobid]
|
||||
grobid-client>=0.8.0
|
||||
1
src/ocr_pipeline.egg-info/top_level.txt
Normal file
1
src/ocr_pipeline.egg-info/top_level.txt
Normal file
@@ -0,0 +1 @@
|
||||
ocr_pipeline
|
||||
12
src/ocr_pipeline/__init__.py
Normal file
12
src/ocr_pipeline/__init__.py
Normal file
@@ -0,0 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
__version__ = "0.1.0"
|
||||
|
||||
from .cli import main
|
||||
from .pipeline import OCRPipeline, PipelineResult
|
||||
|
||||
__all__ = [
|
||||
"OCRPipeline",
|
||||
"PipelineResult",
|
||||
"main",
|
||||
]
|
||||
BIN
src/ocr_pipeline/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/__pycache__/cli.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/__pycache__/cli.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/__pycache__/config.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/__pycache__/config.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/__pycache__/pipeline.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/__pycache__/pipeline.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/__pycache__/watch.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/__pycache__/watch.cpython-313.pyc
Normal file
Binary file not shown.
7
src/ocr_pipeline/citations/__init__.py
Normal file
7
src/ocr_pipeline/citations/__init__.py
Normal file
@@ -0,0 +1,7 @@
|
||||
from .regex_extractor import Citation, extract_citations, format_citations_for_markdown
|
||||
|
||||
__all__ = [
|
||||
"extract_citations",
|
||||
"format_citations_for_markdown",
|
||||
"Citation",
|
||||
]
|
||||
BIN
src/ocr_pipeline/citations/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/citations/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/citations/__pycache__/extractor.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/citations/__pycache__/extractor.cpython-313.pyc
Normal file
Binary file not shown.
Binary file not shown.
191
src/ocr_pipeline/citations/extractor.py
Normal file
191
src/ocr_pipeline/citations/extractor.py
Normal file
@@ -0,0 +1,191 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Citation:
|
||||
type: str
|
||||
value: str
|
||||
start: int
|
||||
end: int
|
||||
confidence: float
|
||||
metadata: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class RegexCitationExtractor:
|
||||
DOI_PATTERN = re.compile(r"\b(?:10\.\d{4,9}/[-._;()/:A-Z0-9]+)\b", re.IGNORECASE)
|
||||
PMID_PATTERN = re.compile(r"\b(?:PMID|PubMed|pubmed)[:/\s]?(\d{1,8})\b", re.IGNORECASE)
|
||||
ARXIV_PATTERN = re.compile(
|
||||
r"\b(?:arXiv|arxiv)[:/\s]?(\d{4}\.\d{4,5}(?:v\d+)?)\b", re.IGNORECASE
|
||||
)
|
||||
PMC_PATTERN = re.compile(r"\b(?:PMC|pmc)[:/\s]?(\d{1,8})\b", re.IGNORECASE)
|
||||
ISBN_PATTERN = re.compile(
|
||||
r"\b(?:ISBN[-:\s]?)?(?:97[89][-\s]?)?\d{1,5}[-\s]?\d{1,7}[-\s]?\d{1,7}[-\s]?[\dX]\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def extract(cls, text: str) -> list[Citation]:
|
||||
citations = []
|
||||
|
||||
for match in cls.DOI_PATTERN.finditer(text):
|
||||
citations.append(
|
||||
Citation(
|
||||
type="DOI",
|
||||
value=match.group(),
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
confidence=0.95,
|
||||
metadata={"url": f"https://doi.org/{match.group()}"},
|
||||
)
|
||||
)
|
||||
|
||||
for match in cls.PMID_PATTERN.finditer(text):
|
||||
citations.append(
|
||||
Citation(
|
||||
type="PMID",
|
||||
value=f"PMID:{match.group(1)}",
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
confidence=0.9,
|
||||
metadata={"url": f"https://pubmed.ncbi.nlm.nih.gov/{match.group(1)}/"},
|
||||
)
|
||||
)
|
||||
|
||||
for match in cls.ARXIV_PATTERN.finditer(text):
|
||||
citations.append(
|
||||
Citation(
|
||||
type="arXiv",
|
||||
value=f"arXiv:{match.group(1)}",
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
confidence=0.9,
|
||||
metadata={"url": f"https://arxiv.org/abs/{match.group(1)}"},
|
||||
)
|
||||
)
|
||||
|
||||
for match in cls.PMC_PATTERN.finditer(text):
|
||||
citations.append(
|
||||
Citation(
|
||||
type="PMC",
|
||||
value=f"PMC:{match.group(1)}",
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
confidence=0.85,
|
||||
metadata={
|
||||
"url": f"https://www.ncbi.nlm.nih.gov/pmc/articles/PMC{match.group(1)}/"
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
for match in cls.ISBN_PATTERN.finditer(text):
|
||||
isbn = re.sub(r"[^\dX]", "", match.group(), flags=re.IGNORECASE)
|
||||
if len(isbn) in (10, 13):
|
||||
citations.append(
|
||||
Citation(
|
||||
type="ISBN",
|
||||
value=isbn,
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
confidence=0.8,
|
||||
metadata={"url": f"https://isbnsearch.org/isbn/{isbn}"},
|
||||
)
|
||||
)
|
||||
|
||||
return citations
|
||||
|
||||
|
||||
class GrobidClient:
|
||||
def __init__(self):
|
||||
self.url = settings.citations.grobid_url
|
||||
self.timeout = settings.citations.grobid_timeout
|
||||
|
||||
def extract_references(self, text: str) -> list[Citation]:
|
||||
if not settings.citations.grobid_enabled:
|
||||
return []
|
||||
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{self.url}/api/processCitationList",
|
||||
data={"citations": text},
|
||||
headers={"Accept": "application/xml"},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
if response.status_code != 200:
|
||||
logger.warning("grobid_request_failed", status=response.status_code)
|
||||
return []
|
||||
|
||||
return self._parse_grobid_xml(response.text)
|
||||
except Exception as e:
|
||||
logger.warning("grobid_extraction_failed", error=str(e))
|
||||
return []
|
||||
|
||||
def _parse_grobid_xml(self, xml_text: str) -> list[Citation]:
|
||||
citations = []
|
||||
try:
|
||||
import xml.etree.ElementTree as ET
|
||||
|
||||
root = ET.fromstring(xml_text)
|
||||
for cit in root.findall(".//{*}citation"):
|
||||
citation_text = ET.tostring(cit, encoding="unicode")
|
||||
doi_match = self.DOI_PATTERN.search(citation_text)
|
||||
if doi_match:
|
||||
citations.append(
|
||||
Citation(
|
||||
type="DOI",
|
||||
value=doi_match.group(),
|
||||
start=0,
|
||||
end=len(citation_text),
|
||||
confidence=0.99,
|
||||
metadata={"source": "grobid", "full_citation": citation_text},
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return citations
|
||||
|
||||
|
||||
def extract_citations(text: str) -> list[Citation]:
|
||||
citations = []
|
||||
|
||||
if settings.citations.regex_enabled:
|
||||
citations.extend(RegexCitationExtractor.extract(text))
|
||||
|
||||
if settings.citations.grobid_enabled:
|
||||
grobid = GrobidClient()
|
||||
citations.extend(grobid.extract_references(text))
|
||||
|
||||
seen = set()
|
||||
unique = []
|
||||
for cit in citations:
|
||||
key = (cit.type, cit.value)
|
||||
if key not in seen:
|
||||
seen.add(key)
|
||||
unique.append(cit)
|
||||
|
||||
return unique
|
||||
|
||||
|
||||
def format_citations_for_markdown(citations: list[Citation]) -> str:
|
||||
if not citations:
|
||||
return "No citations found."
|
||||
|
||||
lines = ["## Citations", ""]
|
||||
for cit in citations:
|
||||
url = cit.metadata.get("url", "") if cit.metadata else ""
|
||||
if url:
|
||||
lines.append(f"- [{cit.type}: {cit.value}]({url})")
|
||||
else:
|
||||
lines.append(f"- {cit.type}: {cit.value}")
|
||||
|
||||
return "\n".join(lines)
|
||||
196
src/ocr_pipeline/citations/regex_extractor.py
Normal file
196
src/ocr_pipeline/citations/regex_extractor.py
Normal file
@@ -0,0 +1,196 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
import httpx
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Citation:
|
||||
raw: str
|
||||
type: str
|
||||
identifier: str
|
||||
confidence: float
|
||||
context: str = ""
|
||||
|
||||
|
||||
class RegexCitationExtractor:
|
||||
DOI_PATTERN = re.compile(
|
||||
r"\b(?:10\.\d{4,9}/[-._;()/:A-Z0-9]+)\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
PMID_PATTERN = re.compile(
|
||||
r"\b(?:PMID|PubMed\s*ID|PubMed)\s*[:\-]?\s*(\d{8})\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
ARXIV_PATTERN = re.compile(
|
||||
r"\b(?:arXiv|arxiv)\s*[:\-]?\s*(\d{4}\.\d{4,5}(?:v\d+)?)\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
ISBN_PATTERN = re.compile(
|
||||
r"\b(?:ISBN[-\s]*(?:13|10)?[-\s]*)?(?:97[89][-\s]*)?\d{1,5}[-\s]*\d{1,7}[-\s]*\d{1,7}[-\s]*[\dX]\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def extract(cls, text: str) -> list[Citation]:
|
||||
citations = []
|
||||
|
||||
for match in cls.DOI_PATTERN.finditer(text):
|
||||
ctx_start = max(0, match.start() - 80)
|
||||
ctx_end = min(len(text), match.end() + 80)
|
||||
citations.append(
|
||||
Citation(
|
||||
raw=match.group(),
|
||||
type="doi",
|
||||
identifier=match.group().lower(),
|
||||
confidence=0.95,
|
||||
context=text[ctx_start:ctx_end].strip(),
|
||||
)
|
||||
)
|
||||
|
||||
for match in cls.PMID_PATTERN.finditer(text):
|
||||
ctx_start = max(0, match.start() - 80)
|
||||
ctx_end = min(len(text), match.end() + 80)
|
||||
citations.append(
|
||||
Citation(
|
||||
raw=match.group(),
|
||||
type="pmid",
|
||||
identifier=f"PMID:{match.group(1)}",
|
||||
confidence=0.9,
|
||||
context=text[ctx_start:ctx_end].strip(),
|
||||
)
|
||||
)
|
||||
|
||||
for match in cls.ARXIV_PATTERN.finditer(text):
|
||||
ctx_start = max(0, match.start() - 80)
|
||||
ctx_end = min(len(text), match.end() + 80)
|
||||
citations.append(
|
||||
Citation(
|
||||
raw=match.group(),
|
||||
type="arxiv",
|
||||
identifier=f"arXiv:{match.group(1)}",
|
||||
confidence=0.9,
|
||||
context=text[ctx_start:ctx_end].strip(),
|
||||
)
|
||||
)
|
||||
|
||||
for match in cls.ISBN_PATTERN.finditer(text):
|
||||
ctx_start = max(0, match.start() - 80)
|
||||
ctx_end = min(len(text), match.end() + 80)
|
||||
isbn = re.sub(r"[-\s]", "", match.group())
|
||||
citations.append(
|
||||
Citation(
|
||||
raw=match.group(),
|
||||
type="isbn",
|
||||
identifier=f"ISBN:{isbn}",
|
||||
confidence=0.7,
|
||||
context=text[ctx_start:ctx_end].strip(),
|
||||
)
|
||||
)
|
||||
|
||||
return citations
|
||||
|
||||
|
||||
class GrobidClient:
|
||||
def __init__(self):
|
||||
self.url = settings.citations.grobid_url
|
||||
self.timeout = settings.citations.grobid_timeout
|
||||
|
||||
def process_text(self, text: str) -> list[Citation]:
|
||||
if not settings.citations.grobid_enabled:
|
||||
return []
|
||||
|
||||
try:
|
||||
with httpx.Client(timeout=self.timeout) as client:
|
||||
response = client.post(
|
||||
f"{self.url}/api/processCitationText",
|
||||
data={"text": text},
|
||||
)
|
||||
response.raise_for_status()
|
||||
return self._parse_grobid_xml(response.text)
|
||||
except Exception as e:
|
||||
logger.warning("grobid_request_failed", error=str(e))
|
||||
return []
|
||||
|
||||
def _parse_grobid_xml(self, xml_text: str) -> list[Citation]:
|
||||
try:
|
||||
import xml.etree.ElementTree as ET
|
||||
|
||||
root = ET.fromstring(xml_text)
|
||||
citations = []
|
||||
|
||||
ns = {"tei": "http://www.tei-c.org/ns/1.0"}
|
||||
for cit in root.findall(".//tei:biblStruct", ns):
|
||||
doi_elem = cit.find(".//tei:idno[@type='DOI']", ns)
|
||||
pmid_elem = cit.find(".//tei:idno[@type='PMID']", ns)
|
||||
|
||||
if doi_elem is not None and doi_elem.text:
|
||||
citations.append(
|
||||
Citation(
|
||||
raw=doi_elem.text,
|
||||
type="doi",
|
||||
identifier=doi_elem.text.lower(),
|
||||
confidence=0.99,
|
||||
)
|
||||
)
|
||||
if pmid_elem is not None and pmid_elem.text:
|
||||
citations.append(
|
||||
Citation(
|
||||
raw=pmid_elem.text,
|
||||
type="pmid",
|
||||
identifier=f"PMID:{pmid_elem.text}",
|
||||
confidence=0.99,
|
||||
)
|
||||
)
|
||||
|
||||
return citations
|
||||
except Exception as e:
|
||||
logger.warning("grobid_xml_parse_failed", error=str(e))
|
||||
return []
|
||||
|
||||
|
||||
def extract_citations(text: str) -> list[Citation]:
|
||||
citations = []
|
||||
|
||||
if settings.citations.regex_enabled:
|
||||
citations.extend(RegexCitationExtractor.extract(text))
|
||||
|
||||
if settings.citations.grobid_enabled:
|
||||
client = GrobidClient()
|
||||
citations.extend(client.process_text(text))
|
||||
|
||||
seen = set()
|
||||
unique = []
|
||||
for cit in citations:
|
||||
key = (cit.type, cit.identifier)
|
||||
if key not in seen:
|
||||
seen.add(key)
|
||||
unique.append(cit)
|
||||
|
||||
return unique
|
||||
|
||||
|
||||
def format_citations_for_markdown(citations: list[Citation]) -> str:
|
||||
if not citations:
|
||||
return "No citations found."
|
||||
|
||||
lines = [
|
||||
"## Citations",
|
||||
"",
|
||||
"| Type | Identifier | Context |",
|
||||
"|------|------------|---------|",
|
||||
]
|
||||
for cit in citations:
|
||||
context = (
|
||||
cit.context[:80].replace("|", "\\|") + "..." if len(cit.context) > 80 else cit.context
|
||||
)
|
||||
lines.append(f"| {cit.type.upper()} | `{cit.identifier}` | {context} |")
|
||||
|
||||
return "\n".join(lines)
|
||||
163
src/ocr_pipeline/cli.py
Normal file
163
src/ocr_pipeline/cli.py
Normal file
@@ -0,0 +1,163 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import typer
|
||||
from rich.console import Console
|
||||
from rich.table import Table
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.pipeline import OCRPipeline
|
||||
from ocr_pipeline.utils.logging import get_logger, setup_logging
|
||||
from ocr_pipeline.utils.migrate import migrate_notebook
|
||||
|
||||
app = typer.Typer(
|
||||
name="ocr-pipeline",
|
||||
help="OCR pipeline for screenshots -> RAG-ready Markdown",
|
||||
add_completion=False,
|
||||
)
|
||||
console = Console()
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@app.callback()
|
||||
def callback(
|
||||
config: Path | None = typer.Option(None, "--config", "-c", help="Config file path"),
|
||||
log_level: str = typer.Option("INFO", "--log-level", "-l", help="Log level"),
|
||||
log_file: Path | None = typer.Option(None, "--log-file", help="Log file path"),
|
||||
):
|
||||
"""OCR Pipeline - Extract text from screenshots for RAG"""
|
||||
setup_logging(log_level, log_file)
|
||||
|
||||
if config:
|
||||
settings.update_from_yaml(config)
|
||||
|
||||
|
||||
@app.command()
|
||||
def run(
|
||||
input_dir: Path | None = typer.Option(None, "--input-dir", "-i", help="Input directory"),
|
||||
output_dir: Path | None = typer.Option(None, "--output-dir", "-o", help="Output directory"),
|
||||
workers: int | None = typer.Option(
|
||||
None, "--workers", "-w", help="Number of worker processes"
|
||||
),
|
||||
engine: str | None = typer.Option(
|
||||
None, "--engine", "-e", help="OCR engine (paddleocr/tesseract/auto)"
|
||||
),
|
||||
):
|
||||
"""Run OCR pipeline on all screenshots"""
|
||||
if input_dir:
|
||||
settings.input.paths = [str(input_dir)]
|
||||
if output_dir:
|
||||
settings.output.base_directory = str(output_dir)
|
||||
if workers:
|
||||
settings.processing.workers = workers
|
||||
if engine:
|
||||
settings.ocr.engine = engine
|
||||
|
||||
pipeline = OCRPipeline()
|
||||
result = pipeline.process_all()
|
||||
|
||||
table = Table(title="Pipeline Results")
|
||||
table.add_column("Metric", style="cyan")
|
||||
table.add_column("Value", style="green")
|
||||
|
||||
table.add_row("Total Images", str(result.total_images))
|
||||
table.add_row("Successful", str(result.successful))
|
||||
table.add_row("Failed", str(result.failed))
|
||||
table.add_row("Chunks Created", str(result.chunks_created))
|
||||
table.add_row("Output Files", str(len(result.output_files)))
|
||||
|
||||
console.print(table)
|
||||
|
||||
if result.errors:
|
||||
console.print("\n[red]Errors:[/red]")
|
||||
for err in result.errors:
|
||||
console.print(f" {err['path']}: {err['error']}")
|
||||
|
||||
if result.failed > 0:
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
@app.command()
|
||||
def watch():
|
||||
"""Watch for new screenshots and process automatically"""
|
||||
pipeline = OCRPipeline()
|
||||
console.print("[green]Starting watch mode...[/green]")
|
||||
console.print("Press Ctrl+C to stop")
|
||||
try:
|
||||
pipeline.watch()
|
||||
except KeyboardInterrupt:
|
||||
console.print("\n[yellow]Stopping watcher...[/yellow]")
|
||||
|
||||
|
||||
@app.command()
|
||||
def reprocess(
|
||||
notebook: Path = typer.Argument(..., help="Path to Jupyter notebook"),
|
||||
):
|
||||
"""Reprocess screenshots from a Jupyter notebook"""
|
||||
pipeline = OCRPipeline()
|
||||
console.print(f"[green]Migrating from notebook: {notebook}[/green]")
|
||||
result = migrate_notebook(notebook, pipeline)
|
||||
|
||||
table = Table(title="Migration Results")
|
||||
table.add_column("Metric", style="cyan")
|
||||
table.add_column("Value", style="green")
|
||||
|
||||
for key, value in result.items():
|
||||
if key != "errors":
|
||||
table.add_row(key.capitalize(), str(value))
|
||||
|
||||
console.print(table)
|
||||
|
||||
if result.get("errors"):
|
||||
console.print("\n[red]Errors:[/red]")
|
||||
for err in result["errors"]:
|
||||
console.print(f" {err['path']}: {err['error']}")
|
||||
|
||||
|
||||
@app.command()
|
||||
def status():
|
||||
"""Show pipeline status and statistics"""
|
||||
pipeline = OCRPipeline()
|
||||
stats = pipeline.get_stats()
|
||||
|
||||
table = Table(title="Pipeline Status")
|
||||
table.add_column("Metric", style="cyan")
|
||||
table.add_column("Value", style="green")
|
||||
|
||||
for key, value in stats.items():
|
||||
table.add_row(key.replace("_", " ").title(), str(value))
|
||||
|
||||
console.print(table)
|
||||
|
||||
|
||||
@app.command()
|
||||
def stats(days: int = typer.Option(30, "--days", "-d", help="Number of days to analyze")):
|
||||
"""Show processing statistics"""
|
||||
pipeline = OCRPipeline()
|
||||
stats = pipeline.get_stats()
|
||||
|
||||
table = Table(title=f"Statistics (last {days} days)")
|
||||
table.add_column("Metric", style="cyan")
|
||||
table.add_column("Value", style="green")
|
||||
|
||||
for key, value in stats.items():
|
||||
table.add_row(key.replace("_", " ").title(), str(value))
|
||||
|
||||
console.print(table)
|
||||
|
||||
|
||||
@app.command()
|
||||
def config():
|
||||
"""Show current configuration"""
|
||||
console.print("[cyan]Current Configuration:[/cyan]")
|
||||
console.print(settings.model_dump_json(indent=2))
|
||||
|
||||
|
||||
def main():
|
||||
app()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
246
src/ocr_pipeline/config.py
Normal file
246
src/ocr_pipeline/config.py
Normal file
@@ -0,0 +1,246 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, ClassVar
|
||||
|
||||
import yaml
|
||||
from pydantic import Field, field_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class InputConfig(BaseSettings):
|
||||
paths: list[str] = Field(default_factory=lambda: ["~/Pictures"])
|
||||
patterns: list[str] = Field(
|
||||
default_factory=lambda: ["SCR-*.png", "*.jpg", "*.jpeg", "*.tiff", "*.bmp"]
|
||||
)
|
||||
recursive: bool = True
|
||||
|
||||
@field_validator("paths", mode="before")
|
||||
@classmethod
|
||||
def expand_paths(cls, v: list[str]) -> list[str]:
|
||||
return [os.path.expanduser(p) for p in v]
|
||||
|
||||
|
||||
class PreprocessConfig(BaseSettings):
|
||||
deskew: bool = True
|
||||
denoise: bool = True
|
||||
clahe: bool = True
|
||||
adaptive_threshold: bool = True
|
||||
remove_lines: bool = True
|
||||
max_dimension: int = 4096
|
||||
|
||||
|
||||
class OCRConfig(BaseSettings):
|
||||
engine: str = "paddleocr" # paddleocr | tesseract | auto
|
||||
languages: list[str] = Field(default_factory=lambda: ["en", "latin"])
|
||||
use_gpu: bool = False
|
||||
use_angle_cls: bool = True
|
||||
det_db_thresh: float = 0.3
|
||||
det_db_box_thresh: float = 0.6
|
||||
det_db_unclip_ratio: float = 1.5
|
||||
rec_batch_num: int = 6
|
||||
cpu_threads: int = 4
|
||||
preprocess: PreprocessConfig = Field(default_factory=PreprocessConfig)
|
||||
|
||||
|
||||
class ProcessingConfig(BaseSettings):
|
||||
workers: int = 4
|
||||
batch_size: int = 10
|
||||
retry_attempts: int = 2
|
||||
timeout_per_image: int = 120
|
||||
skip_existing: bool = True
|
||||
|
||||
|
||||
class DetectorConfig(BaseSettings):
|
||||
enabled: bool = True
|
||||
model: str = ""
|
||||
confidence_threshold: float = 0.7
|
||||
|
||||
|
||||
class FigureDetectorConfig(DetectorConfig):
|
||||
model: str = "lp://PubLayNet/faster_rcnn_R_50_FPN_3x/config"
|
||||
confidence_threshold: float = 0.7
|
||||
|
||||
|
||||
class TableDetectorConfig(DetectorConfig):
|
||||
model: str = "microsoft/table-transformer-detection"
|
||||
confidence_threshold: float = 0.7
|
||||
|
||||
|
||||
class CaptionDetectorConfig(DetectorConfig):
|
||||
model: str = "lp://PubLayNet/faster_rcnn_R_50_FPN_3x/config"
|
||||
confidence_threshold: float = 0.5
|
||||
|
||||
|
||||
class DetectorsConfig(BaseSettings):
|
||||
figures: FigureDetectorConfig = Field(default_factory=FigureDetectorConfig)
|
||||
tables: TableDetectorConfig = Field(default_factory=TableDetectorConfig)
|
||||
captions: CaptionDetectorConfig = Field(default_factory=CaptionDetectorConfig)
|
||||
|
||||
|
||||
class EntitiesConfig(BaseSettings):
|
||||
enabled: bool = True
|
||||
model: str = "en_core_sci_lg"
|
||||
types: list[str] = Field(
|
||||
default_factory=lambda: [
|
||||
"GENE",
|
||||
"PROTEIN",
|
||||
"CHEMICAL",
|
||||
"SPECIES",
|
||||
"DISEASE",
|
||||
"CELL_LINE",
|
||||
"ORGANISM",
|
||||
"CELL_TYPE",
|
||||
]
|
||||
)
|
||||
merge_entities: bool = True
|
||||
|
||||
|
||||
class CitationsConfig(BaseSettings):
|
||||
regex_enabled: bool = True
|
||||
grobid_enabled: bool = False
|
||||
grobid_url: str = "http://localhost:8070"
|
||||
grobid_timeout: int = 30
|
||||
|
||||
|
||||
class ChunkingConfig(BaseSettings):
|
||||
chunk_size: int = 1000
|
||||
chunk_overlap: int = 200
|
||||
separators: list[str] = Field(default_factory=lambda: ["\n\n", "\n", ". ", " ", ""])
|
||||
keep_separator: bool = True
|
||||
|
||||
|
||||
class FrontmatterConfig(BaseSettings):
|
||||
fields: list[str] = Field(
|
||||
default_factory=lambda: [
|
||||
"source_path",
|
||||
"source_hash",
|
||||
"timestamp",
|
||||
"ocr_engine",
|
||||
"ocr_confidence_mean",
|
||||
"language",
|
||||
"detected_entities",
|
||||
"has_figures",
|
||||
"has_tables",
|
||||
"citations_found",
|
||||
"chunk_index",
|
||||
"total_chunks",
|
||||
]
|
||||
)
|
||||
include_raw_text: bool = False
|
||||
|
||||
|
||||
class OutputConfig(BaseSettings):
|
||||
base_directory: str = "./data/ocr_output"
|
||||
format: str = "markdown"
|
||||
organize_by: str = "date_run" # date_run | source_dir | flat
|
||||
frontmatter: FrontmatterConfig = Field(default_factory=FrontmatterConfig)
|
||||
|
||||
|
||||
class WatchConfig(BaseSettings):
|
||||
enabled: bool = True
|
||||
debounce_seconds: int = 5
|
||||
ignore_patterns: list[str] = Field(
|
||||
default_factory=lambda: [
|
||||
".DS_Store",
|
||||
"*.tmp",
|
||||
"*.partial",
|
||||
"*.crdownload",
|
||||
]
|
||||
)
|
||||
db_path: str = "./data/processed_files.db"
|
||||
poll_interval: float = 1.0
|
||||
|
||||
|
||||
class _BaseSettings(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
env_nested_delimiter="__",
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
|
||||
def _load_yaml_config(path: str = "config.yaml") -> dict[str, Any]:
|
||||
"""Load YAML config file if it exists."""
|
||||
config_path = Path(path)
|
||||
if not config_path.is_absolute():
|
||||
config_path = Path.cwd() / config_path
|
||||
if config_path.exists():
|
||||
with open(config_path) as f:
|
||||
return yaml.safe_load(f) or {}
|
||||
return {}
|
||||
|
||||
|
||||
def _build_settings_from_yaml(yaml_dict: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Convert YAML dict to properly typed nested models for Settings."""
|
||||
if not yaml_dict:
|
||||
return {}
|
||||
|
||||
result = {}
|
||||
|
||||
# Map YAML keys to model classes
|
||||
model_map = {
|
||||
"input": InputConfig,
|
||||
"ocr": OCRConfig,
|
||||
"processing": ProcessingConfig,
|
||||
"detectors": DetectorsConfig,
|
||||
"entities": EntitiesConfig,
|
||||
"citations": CitationsConfig,
|
||||
"chunking": ChunkingConfig,
|
||||
"output": OutputConfig,
|
||||
"watch": WatchConfig,
|
||||
}
|
||||
|
||||
for key, model_class in model_map.items():
|
||||
if key in yaml_dict:
|
||||
data = yaml_dict[key]
|
||||
# Special handling for output config to transform frontmatter list -> dict
|
||||
if key == "output" and "frontmatter" in data and isinstance(data["frontmatter"], list):
|
||||
data = {**data, "frontmatter": {"fields": data["frontmatter"]}}
|
||||
result[key] = model_class(**data)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class Settings(_BaseSettings):
|
||||
input: InputConfig = Field(default_factory=InputConfig)
|
||||
ocr: OCRConfig = Field(default_factory=OCRConfig)
|
||||
processing: ProcessingConfig = Field(default_factory=ProcessingConfig)
|
||||
detectors: DetectorsConfig = Field(default_factory=DetectorsConfig)
|
||||
entities: EntitiesConfig = Field(default_factory=EntitiesConfig)
|
||||
citations: CitationsConfig = Field(default_factory=CitationsConfig)
|
||||
chunking: ChunkingConfig = Field(default_factory=ChunkingConfig)
|
||||
output: OutputConfig = Field(default_factory=OutputConfig)
|
||||
watch: WatchConfig = Field(default_factory=WatchConfig)
|
||||
|
||||
# Store YAML config for manual override
|
||||
_yaml_config: ClassVar[dict[str, Any]] = {}
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
if not Settings._yaml_config:
|
||||
Settings._yaml_config = _load_yaml_config()
|
||||
|
||||
# Build nested models from YAML
|
||||
yaml_models = _build_settings_from_yaml(Settings._yaml_config)
|
||||
super().__init__(**yaml_models, **kwargs)
|
||||
|
||||
def update_from_yaml(self, path: str) -> None:
|
||||
"""Reload settings from a YAML file."""
|
||||
Settings._yaml_config = _load_yaml_config(path)
|
||||
# Re-initialize with new config
|
||||
yaml_models = _build_settings_from_yaml(Settings._yaml_config)
|
||||
for key, value in yaml_models.items():
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
|
||||
@property
|
||||
def output_dir(self) -> Path:
|
||||
return Path(self.output.base_directory).expanduser().resolve()
|
||||
|
||||
@property
|
||||
def watch_db_path(self) -> Path:
|
||||
return Path(self.watch.db_path).expanduser().resolve()
|
||||
|
||||
|
||||
settings = Settings()
|
||||
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
|
||||
22
src/ocr_pipeline/ocr/__init__.py
Normal file
22
src/ocr_pipeline/ocr/__init__.py
Normal file
@@ -0,0 +1,22 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from .engine import (
|
||||
ImagePreprocessor,
|
||||
OCREngine,
|
||||
OCRResult,
|
||||
PaddleOCREngine,
|
||||
TesseractEngine,
|
||||
compute_image_hash,
|
||||
)
|
||||
from .parallel import ParallelProcessor, ProcessedImage
|
||||
|
||||
__all__ = [
|
||||
"OCREngine",
|
||||
"OCRResult",
|
||||
"ImagePreprocessor",
|
||||
"compute_image_hash",
|
||||
"PaddleOCREngine",
|
||||
"TesseractEngine",
|
||||
"ParallelProcessor",
|
||||
"ProcessedImage",
|
||||
]
|
||||
BIN
src/ocr_pipeline/ocr/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/ocr/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/ocr/__pycache__/engine.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/ocr/__pycache__/engine.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/ocr/__pycache__/parallel.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/ocr/__pycache__/parallel.cpython-313.pyc
Normal file
Binary file not shown.
236
src/ocr_pipeline/ocr/engine.py
Normal file
236
src/ocr_pipeline/ocr/engine.py
Normal file
@@ -0,0 +1,236 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pytesseract
|
||||
from PIL import Image
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Language code mapping: config names -> engine-specific codes
|
||||
TESSERACT_LANG_MAP = {
|
||||
"en": "eng",
|
||||
"latin": "lat",
|
||||
"de": "deu",
|
||||
"fr": "fra",
|
||||
"es": "spa",
|
||||
"it": "ita",
|
||||
"pt": "por",
|
||||
"zh": "chi_sim",
|
||||
"ja": "jpn",
|
||||
"ko": "kor",
|
||||
}
|
||||
|
||||
PADDLEOCR_LANG_MAP = {
|
||||
"en": "en",
|
||||
"latin": "latin",
|
||||
"de": "german",
|
||||
"fr": "french",
|
||||
"es": "spanish",
|
||||
"it": "italian",
|
||||
"pt": "portuguese",
|
||||
"zh": "ch",
|
||||
"ja": "japan",
|
||||
"ko": "korean",
|
||||
}
|
||||
|
||||
|
||||
def map_languages(languages: list[str], engine: str) -> str:
|
||||
"""Map config language names to engine-specific codes."""
|
||||
if engine == "tesseract":
|
||||
mapped = [TESSERACT_LANG_MAP.get(lang, lang) for lang in languages]
|
||||
return "+".join(mapped)
|
||||
else: # paddleocr
|
||||
# PaddleOCR uses primary language only for detection
|
||||
primary = languages[0] if languages else "en"
|
||||
return PADDLEOCR_LANG_MAP.get(primary, primary)
|
||||
|
||||
|
||||
@dataclass
|
||||
class OCRResult:
|
||||
text: str
|
||||
confidence: float
|
||||
language: str
|
||||
engine: str
|
||||
boxes: list[dict[str, Any]]
|
||||
raw_result: Any
|
||||
processing_time: float
|
||||
|
||||
|
||||
class PaddleOCREngine:
|
||||
def __init__(self):
|
||||
self._ocr = None
|
||||
self._initialized = False
|
||||
|
||||
def _init(self):
|
||||
if self._initialized:
|
||||
return
|
||||
try:
|
||||
from paddleocr import PaddleOCR
|
||||
|
||||
paddle_lang = map_languages(settings.ocr.languages, "paddleocr")
|
||||
|
||||
self._ocr = PaddleOCR(
|
||||
use_angle_cls=settings.ocr.use_angle_cls,
|
||||
lang=paddle_lang,
|
||||
use_gpu=settings.ocr.use_gpu,
|
||||
show_log=False,
|
||||
det_db_thresh=settings.ocr.det_db_thresh,
|
||||
det_db_box_thresh=settings.ocr.det_db_box_thresh,
|
||||
det_db_unclip_ratio=settings.ocr.det_db_unclip_ratio,
|
||||
rec_batch_num=settings.ocr.rec_batch_num,
|
||||
cpu_threads=settings.ocr.cpu_threads,
|
||||
)
|
||||
self._initialized = True
|
||||
logger.info("paddleocr_initialized", language=paddle_lang)
|
||||
except Exception as e:
|
||||
logger.error("paddleocr_init_failed", error=str(e))
|
||||
raise
|
||||
|
||||
def process(self, image: np.ndarray) -> OCRResult:
|
||||
self._init()
|
||||
start = time.time()
|
||||
|
||||
result = self._ocr.ocr(image, cls=True)
|
||||
processing_time = time.time() - start
|
||||
|
||||
if not result or not result[0]:
|
||||
return OCRResult(
|
||||
text="",
|
||||
confidence=0.0,
|
||||
language=map_languages(settings.ocr.languages, "paddleocr"),
|
||||
engine="paddleocr",
|
||||
boxes=[],
|
||||
raw_result=result,
|
||||
processing_time=processing_time,
|
||||
)
|
||||
|
||||
texts = []
|
||||
confidences = []
|
||||
boxes = []
|
||||
|
||||
for line in result[0]:
|
||||
box = line[0]
|
||||
text = line[1][0]
|
||||
conf = line[1][1]
|
||||
|
||||
texts.append(text)
|
||||
confidences.append(conf)
|
||||
boxes.append(
|
||||
{
|
||||
"bbox": box,
|
||||
"text": text,
|
||||
"confidence": conf,
|
||||
}
|
||||
)
|
||||
|
||||
full_text = "\n".join(texts)
|
||||
mean_conf = sum(confidences) / len(confidences) if confidences else 0.0
|
||||
|
||||
return OCRResult(
|
||||
text=full_text,
|
||||
confidence=mean_conf,
|
||||
language=map_languages(settings.ocr.languages, "paddleocr"),
|
||||
engine="paddleocr",
|
||||
boxes=boxes,
|
||||
raw_result=result,
|
||||
processing_time=processing_time,
|
||||
)
|
||||
|
||||
|
||||
class TesseractEngine:
|
||||
def process(self, image: np.ndarray) -> OCRResult:
|
||||
start = time.time()
|
||||
|
||||
pil_image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB))
|
||||
|
||||
tess_lang = map_languages(settings.ocr.languages, "tesseract")
|
||||
config = f"--oem 3 --psm 6 -l {tess_lang}"
|
||||
data = pytesseract.image_to_data(
|
||||
pil_image, config=config, output_type=pytesseract.Output.DICT
|
||||
)
|
||||
|
||||
texts = []
|
||||
confidences = []
|
||||
boxes = []
|
||||
|
||||
for i in range(len(data["text"])):
|
||||
text = data["text"][i].strip()
|
||||
conf = data["conf"][i]
|
||||
if text and conf > 0:
|
||||
texts.append(text)
|
||||
confidences.append(conf / 100.0)
|
||||
boxes.append(
|
||||
{
|
||||
"bbox": [
|
||||
[data["left"][i], data["top"][i]],
|
||||
[data["left"][i] + data["width"][i], data["top"][i]],
|
||||
[
|
||||
data["left"][i] + data["width"][i],
|
||||
data["top"][i] + data["height"][i],
|
||||
],
|
||||
[data["left"][i], data["top"][i] + data["height"][i]],
|
||||
],
|
||||
"text": text,
|
||||
"confidence": conf / 100.0,
|
||||
}
|
||||
)
|
||||
|
||||
full_text = " ".join(texts)
|
||||
mean_conf = sum(confidences) / len(confidences) if confidences else 0.0
|
||||
|
||||
return OCRResult(
|
||||
text=full_text,
|
||||
confidence=mean_conf,
|
||||
language=map_languages(settings.ocr.languages, "tesseract"),
|
||||
engine="tesseract",
|
||||
boxes=boxes,
|
||||
raw_result=data,
|
||||
processing_time=time.time() - start,
|
||||
)
|
||||
|
||||
|
||||
class OCREngine:
|
||||
def __init__(self):
|
||||
self.paddle = PaddleOCREngine()
|
||||
self.tesseract = TesseractEngine()
|
||||
self.engine_preference = settings.ocr.engine
|
||||
|
||||
def process(self, image: np.ndarray) -> OCRResult:
|
||||
"""Process a pre-loaded image array."""
|
||||
if self.engine_preference == "paddleocr":
|
||||
try:
|
||||
return self.paddle.process(image)
|
||||
except Exception as e:
|
||||
logger.warning("paddleocr_failed_fallback", error=str(e))
|
||||
return self.tesseract.process(image)
|
||||
|
||||
elif self.engine_preference == "tesseract":
|
||||
return self.tesseract.process(image)
|
||||
|
||||
else:
|
||||
try:
|
||||
return self.paddle.process(image)
|
||||
except Exception as e:
|
||||
logger.warning("paddleocr_failed_fallback", error=str(e))
|
||||
return self.tesseract.process(image)
|
||||
|
||||
def process_file(self, image_path: Path) -> OCRResult:
|
||||
"""Load and process an image file."""
|
||||
logger.info("processing_image", path=str(image_path), engine=self.engine_preference)
|
||||
|
||||
image = cv2.imread(str(image_path))
|
||||
if image is None:
|
||||
raise ValueError(f"Failed to load image: {image_path}")
|
||||
return self.process(image)
|
||||
|
||||
# Alias for backward compatibility
|
||||
recognize = process
|
||||
147
src/ocr_pipeline/ocr/parallel.py
Normal file
147
src/ocr_pipeline/ocr/parallel.py
Normal file
@@ -0,0 +1,147 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from concurrent.futures import ProcessPoolExecutor, as_completed
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import ProgressTracker, get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProcessedImage:
|
||||
path: Path
|
||||
image_hash: str
|
||||
ocr_result: Any
|
||||
preprocessing_time: float
|
||||
ocr_time: float
|
||||
success: bool
|
||||
error: str | None = None
|
||||
|
||||
|
||||
def _process_single_image(args: tuple[Path, dict]) -> ProcessedImage:
|
||||
"""Worker function for multiprocessing - must be at module level."""
|
||||
path, config_dict = args
|
||||
from ocr_pipeline.ocr.engine import ImagePreprocessor, OCREngine, compute_image_hash
|
||||
|
||||
preprocessor = ImagePreprocessor()
|
||||
ocr_engine = OCREngine()
|
||||
|
||||
try:
|
||||
image_hash = compute_image_hash(path)
|
||||
prep_start = time.time()
|
||||
image = preprocessor.preprocess(path)
|
||||
preprocessing_time = time.time() - prep_start
|
||||
|
||||
ocr_start = time.time()
|
||||
result = ocr_engine.process(image)
|
||||
ocr_time = time.time() - ocr_start
|
||||
|
||||
return ProcessedImage(
|
||||
path=path,
|
||||
image_hash=image_hash,
|
||||
ocr_result=result,
|
||||
preprocessing_time=preprocessing_time,
|
||||
ocr_time=ocr_time,
|
||||
success=True,
|
||||
)
|
||||
except Exception as e:
|
||||
return ProcessedImage(
|
||||
path=path,
|
||||
image_hash="",
|
||||
ocr_result=None,
|
||||
preprocessing_time=0,
|
||||
ocr_time=0,
|
||||
success=False,
|
||||
error=str(e),
|
||||
)
|
||||
|
||||
|
||||
class ParallelProcessor:
|
||||
def __init__(self, workers: int | None = None):
|
||||
self.workers = workers or settings.processing.workers
|
||||
self.batch_size = settings.processing.batch_size
|
||||
self.retry_attempts = settings.processing.retry_attempts
|
||||
self.timeout = settings.processing.timeout_per_image
|
||||
|
||||
def find_images(self) -> list[Path]:
|
||||
all_images = []
|
||||
for base_path_str in settings.input.paths:
|
||||
base_path = Path(base_path_str).expanduser()
|
||||
if not base_path.exists():
|
||||
logger.warning("input_path_not_found", path=str(base_path))
|
||||
continue
|
||||
|
||||
for pattern in settings.input.patterns:
|
||||
if settings.input.recursive:
|
||||
images = list(base_path.rglob(pattern))
|
||||
else:
|
||||
images = list(base_path.glob(pattern))
|
||||
all_images.extend(images)
|
||||
|
||||
unique_images = []
|
||||
seen = set()
|
||||
for img in all_images:
|
||||
if img.is_file() and img.suffix.lower() in (".png", ".jpg", ".jpeg", ".tiff", ".bmp"):
|
||||
if img not in seen:
|
||||
seen.add(img)
|
||||
unique_images.append(img)
|
||||
|
||||
logger.info("found_images", count=len(unique_images))
|
||||
return sorted(unique_images)
|
||||
|
||||
def process(self, image_paths: list[Path] | None = None) -> list[ProcessedImage]:
|
||||
if image_paths is None:
|
||||
image_paths = self.find_images()
|
||||
|
||||
if not image_paths:
|
||||
logger.warning("no_images_to_process")
|
||||
return []
|
||||
|
||||
results = []
|
||||
failed = []
|
||||
|
||||
with ProgressTracker(len(image_paths), "OCR Processing") as tracker:
|
||||
with ProcessPoolExecutor(max_workers=self.workers) as executor:
|
||||
future_to_path = {
|
||||
executor.submit(_process_single_image, (path, {})): path for path in image_paths
|
||||
}
|
||||
|
||||
for future in as_completed(future_to_path):
|
||||
path = future_to_path[future]
|
||||
try:
|
||||
result = future.result(timeout=self.timeout)
|
||||
if not result.success and self.retry_attempts > 0:
|
||||
for attempt in range(self.retry_attempts):
|
||||
logger.warning("retrying", path=str(path), attempt=attempt + 1)
|
||||
retry_result = _process_single_image((path, {}))
|
||||
if retry_result.success:
|
||||
result = retry_result
|
||||
break
|
||||
results.append(result)
|
||||
tracker.update(success=result.success)
|
||||
tracker.log_result(path, result.success, result.error)
|
||||
except Exception as e:
|
||||
logger.error("processing_failed", path=str(path), error=str(e))
|
||||
results.append(
|
||||
ProcessedImage(
|
||||
path=path,
|
||||
image_hash="",
|
||||
ocr_result=None,
|
||||
preprocessing_time=0,
|
||||
ocr_time=0,
|
||||
success=False,
|
||||
error=str(e),
|
||||
)
|
||||
)
|
||||
tracker.update(success=False)
|
||||
tracker.log_result(path, False, str(e))
|
||||
|
||||
successful = [r for r in results if r.success]
|
||||
failed = [r for r in results if not r.success]
|
||||
logger.info("processing_complete", successful=len(successful), failed=len(failed))
|
||||
return results
|
||||
8
src/ocr_pipeline/output/__init__.py
Normal file
8
src/ocr_pipeline/output/__init__.py
Normal file
@@ -0,0 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from .markdown import MarkdownWriter, OutputMetadata
|
||||
|
||||
__all__ = [
|
||||
"MarkdownWriter",
|
||||
"OutputMetadata",
|
||||
]
|
||||
BIN
src/ocr_pipeline/output/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/output/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/output/__pycache__/markdown.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/output/__pycache__/markdown.cpython-313.pyc
Normal file
Binary file not shown.
143
src/ocr_pipeline/output/markdown.py
Normal file
143
src/ocr_pipeline/output/markdown.py
Normal file
@@ -0,0 +1,143 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class MarkdownWriter:
|
||||
def __init__(self):
|
||||
self.base_dir = Path(settings.output.base_directory).expanduser().resolve()
|
||||
self.organize_by = settings.output.organize_by
|
||||
self.frontmatter_fields = settings.output.frontmatter.fields
|
||||
|
||||
def _get_output_dir(self, timestamp: str) -> Path:
|
||||
dt = datetime.datetime.fromisoformat(timestamp.replace("Z", "+00:00"))
|
||||
|
||||
if self.organize_by == "date_run":
|
||||
date_dir = self.base_dir / dt.strftime("%Y-%m-%d")
|
||||
run_dirs = sorted(date_dir.glob("run_*"))
|
||||
run_num = len(run_dirs) + 1
|
||||
return date_dir / f"run_{run_num:03d}"
|
||||
elif self.organize_by == "source_dir":
|
||||
return self.base_dir / "by_source"
|
||||
else:
|
||||
return self.base_dir
|
||||
|
||||
def write_chunk(
|
||||
self,
|
||||
content: str,
|
||||
metadata: dict[str, Any],
|
||||
figures: list[Any] = None,
|
||||
tables: list[Any] = None,
|
||||
entities: Any = None,
|
||||
citations: list[Any] = None,
|
||||
) -> Path:
|
||||
output_dir = self._get_output_dir(metadata.get("timestamp", ""))
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
source_path = Path(metadata.get("source_path", "unknown"))
|
||||
safe_name = "".join(c for c in source_path.stem if c.isalnum() or c in "-_")[:80]
|
||||
chunk_idx = metadata.get("chunk_index", 0)
|
||||
total_chunks = metadata.get("total_chunks", 1)
|
||||
|
||||
if total_chunks > 1:
|
||||
filename = f"{safe_name}_chunk_{chunk_idx:03d}.md"
|
||||
else:
|
||||
filename = f"{safe_name}.md"
|
||||
|
||||
output_path = output_dir / filename
|
||||
|
||||
frontmatter = self._build_frontmatter(metadata)
|
||||
body = self._build_body(content, figures, tables, entities, citations)
|
||||
|
||||
with open(output_path, "w") as f:
|
||||
f.write("---\n")
|
||||
yaml.dump(frontmatter, f, sort_keys=False, allow_unicode=True)
|
||||
f.write("---\n\n")
|
||||
f.write(body)
|
||||
|
||||
logger.debug("markdown_written", path=str(output_path))
|
||||
return output_path
|
||||
|
||||
def _build_frontmatter(self, metadata: dict[str, Any]) -> dict[str, Any]:
|
||||
fm = {}
|
||||
for field in self.frontmatter_fields:
|
||||
if field in metadata:
|
||||
fm[field] = metadata[field]
|
||||
return fm
|
||||
|
||||
def _build_body(
|
||||
self,
|
||||
content: str,
|
||||
figures: list[Any] = None,
|
||||
tables: list[Any] = None,
|
||||
entities: Any = None,
|
||||
citations: list[Any] = None,
|
||||
) -> str:
|
||||
parts = []
|
||||
|
||||
source_name = Path(content.split("\n")[0]).name if content else "Unknown"
|
||||
parts.append(f"# Screenshot: {source_name}")
|
||||
parts.append("")
|
||||
|
||||
if content:
|
||||
parts.append("## Extracted Text")
|
||||
parts.append("")
|
||||
parts.append(content.strip())
|
||||
parts.append("")
|
||||
|
||||
if figures:
|
||||
parts.append("## Figures Detected")
|
||||
parts.append("")
|
||||
for i, fig in enumerate(figures):
|
||||
parts.append(f"### Figure {i + 1}")
|
||||
parts.append(f"- **Bounding Box**: {fig.bbox}")
|
||||
parts.append(f"- **Confidence**: {fig.confidence:.2f}")
|
||||
if fig.caption:
|
||||
parts.append(f"- **Caption**: {fig.caption}")
|
||||
parts.append("")
|
||||
|
||||
if tables:
|
||||
parts.append("## Tables Detected")
|
||||
parts.append("")
|
||||
for i, table in enumerate(tables):
|
||||
parts.append(f"### Table {i + 1}")
|
||||
parts.append(f"- **Bounding Box**: {table.bbox}")
|
||||
parts.append(f"- **Confidence**: {table.confidence:.2f}")
|
||||
if table.markdown:
|
||||
parts.append("")
|
||||
parts.append("#### Table Content (Markdown)")
|
||||
parts.append("")
|
||||
parts.append(table.markdown)
|
||||
parts.append("")
|
||||
|
||||
if entities and entities.entities:
|
||||
parts.append("## Detected Entities")
|
||||
parts.append("")
|
||||
parts.append("| Entity | Type | Context |")
|
||||
parts.append("|--------|------|---------|")
|
||||
for ent in entities.entities[:50]:
|
||||
context = ent.context[:80].replace("|", "\\|") + (
|
||||
"..." if len(ent.context) > 80 else ""
|
||||
)
|
||||
parts.append(f"| {ent.text} | {ent.label} | {context} |")
|
||||
parts.append("")
|
||||
|
||||
if citations:
|
||||
parts.append("## Citations Found")
|
||||
parts.append("")
|
||||
parts.append("| Type | Identifier |")
|
||||
parts.append("|------|------------|")
|
||||
for cit in citations:
|
||||
parts.append(f"| {cit.type.upper()} | `{cit.identifier}` |")
|
||||
parts.append("")
|
||||
|
||||
return "\n".join(parts)
|
||||
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,
|
||||
}
|
||||
17
src/ocr_pipeline/postprocess/__init__.py
Normal file
17
src/ocr_pipeline/postprocess/__init__.py
Normal file
@@ -0,0 +1,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from .chunk import SemanticChunker, TextChunk, chunk_text
|
||||
from .clean import CleanResult, clean_text
|
||||
from .entities import Entity, SciSpacyNER, extract_entities, summarize_entities
|
||||
|
||||
__all__ = [
|
||||
"clean_text",
|
||||
"CleanResult",
|
||||
"chunk_text",
|
||||
"TextChunk",
|
||||
"SemanticChunker",
|
||||
"extract_entities",
|
||||
"summarize_entities",
|
||||
"Entity",
|
||||
"SciSpacyNER",
|
||||
]
|
||||
Binary file not shown.
BIN
src/ocr_pipeline/postprocess/__pycache__/chunk.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/postprocess/__pycache__/chunk.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/postprocess/__pycache__/clean.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/postprocess/__pycache__/clean.cpython-313.pyc
Normal file
Binary file not shown.
Binary file not shown.
61
src/ocr_pipeline/postprocess/chunk.py
Normal file
61
src/ocr_pipeline/postprocess/chunk.py
Normal file
@@ -0,0 +1,61 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from langchain.text_splitter import RecursiveCharacterTextSplitter
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextChunk:
|
||||
content: str
|
||||
chunk_index: int
|
||||
start_char: int
|
||||
end_char: int
|
||||
metadata: dict[str, Any]
|
||||
|
||||
|
||||
class SemanticChunker:
|
||||
def __init__(self):
|
||||
self.splitter = RecursiveCharacterTextSplitter(
|
||||
chunk_size=settings.chunking.chunk_size,
|
||||
chunk_overlap=settings.chunking.chunk_overlap,
|
||||
separators=settings.chunking.separators,
|
||||
keep_separator=settings.chunking.keep_separator,
|
||||
length_function=len,
|
||||
)
|
||||
|
||||
def chunk(self, text: str, base_metadata: dict[str, Any] | None = None) -> list[TextChunk]:
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
base_metadata = base_metadata or {}
|
||||
docs = self.splitter.create_documents([text], metadatas=[base_metadata])
|
||||
|
||||
chunks = []
|
||||
for i, doc in enumerate(docs):
|
||||
start = text.find(doc.page_content)
|
||||
if start == -1:
|
||||
start = 0
|
||||
chunks.append(
|
||||
TextChunk(
|
||||
content=doc.page_content,
|
||||
chunk_index=i,
|
||||
start_char=start,
|
||||
end_char=start + len(doc.page_content),
|
||||
metadata={**doc.metadata, "chunk_index": i, "total_chunks": len(docs)},
|
||||
)
|
||||
)
|
||||
|
||||
logger.debug("chunked_text", chunks=len(chunks), original_length=len(text))
|
||||
return chunks
|
||||
|
||||
|
||||
def chunk_text(text: str, metadata: dict[str, Any] | None = None) -> list[TextChunk]:
|
||||
chunker = SemanticChunker()
|
||||
return chunker.chunk(text, metadata)
|
||||
163
src/ocr_pipeline/postprocess/clean.py
Normal file
163
src/ocr_pipeline/postprocess/clean.py
Normal file
@@ -0,0 +1,163 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import unicodedata
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class CleanResult:
|
||||
text: str
|
||||
original_length: int
|
||||
cleaned_length: int
|
||||
operations: list[str]
|
||||
|
||||
|
||||
def fix_encoding(text: str) -> str:
|
||||
try:
|
||||
return text.encode("latin-1").decode("utf-8")
|
||||
except UnicodeError:
|
||||
return text
|
||||
|
||||
|
||||
def normalize_unicode(text: str) -> str:
|
||||
return unicodedata.normalize("NFC", text)
|
||||
|
||||
|
||||
def fix_ligatures(text: str) -> str:
|
||||
ligatures = {
|
||||
"fi": "fi",
|
||||
"fl": "fl",
|
||||
"ff": "ff",
|
||||
"ffi": "ffi",
|
||||
"ffl": "ffl",
|
||||
"ſt": "ft",
|
||||
"st": "st",
|
||||
"Ꜳ": "AA",
|
||||
"ꜳ": "aa",
|
||||
"Ꜵ": "AO",
|
||||
"ꜵ": "ao",
|
||||
"Ꜷ": "AU",
|
||||
"ꜷ": "au",
|
||||
"Ꜹ": "AV",
|
||||
"ꜹ": "av",
|
||||
"Ꜻ": "AV",
|
||||
"ꜻ": "av",
|
||||
"Ꜽ": "AY",
|
||||
"ꜽ": "ay",
|
||||
}
|
||||
for lig, repl in ligatures.items():
|
||||
text = text.replace(lig, repl)
|
||||
return text
|
||||
|
||||
|
||||
def remove_control_chars(text: str) -> str:
|
||||
return "".join(ch for ch in text if unicodedata.category(ch)[0] != "C" or ch in "\n\t")
|
||||
|
||||
|
||||
def fix_hyphenation(text: str) -> str:
|
||||
return re.sub(r"(\w+)-\n(\w+)", r"\1\2", text)
|
||||
|
||||
|
||||
def normalize_whitespace(text: str) -> str:
|
||||
text = re.sub(r"[ \t]+", " ", text)
|
||||
text = re.sub(r"\n{3,}", "\n\n", text)
|
||||
return text.strip()
|
||||
|
||||
|
||||
def remove_ocr_artifacts(text: str) -> str:
|
||||
patterns = [
|
||||
r"^\s*[\|┃║]\s*$",
|
||||
r"^\s*[─━]\s*$",
|
||||
r"^\s*[┌┐└┘├┤┬┴┼]\s*$",
|
||||
r"(.)\1{10,}",
|
||||
r"^\s*[•·▪▫◦‣⁃]\s*$",
|
||||
]
|
||||
lines = text.split("\n")
|
||||
cleaned = []
|
||||
for line in lines:
|
||||
if any(re.match(p, line.strip()) for p in patterns):
|
||||
continue
|
||||
cleaned.append(line)
|
||||
return "\n".join(cleaned)
|
||||
|
||||
|
||||
def fix_scientific_notation(text: str) -> str:
|
||||
text = re.sub(r"(\d)\s*[×xX]\s*10\s*[\^]?\s*([+-]?\d+)", r"\1×10^\2", text)
|
||||
text = re.sub(r"(\d)\s*[×xX]\s*10\s*([+-]?\d+)", r"\1×10^\2", text)
|
||||
return text
|
||||
|
||||
|
||||
def normalize_units(text: str) -> str:
|
||||
unit_fixes = {
|
||||
r"(\d)\s*u[mM]\b": r"\1 µm",
|
||||
r"(\d)\s*[µu]g\b": r"\1 µg",
|
||||
r"(\d)\s*[mn]g\b": r"\1 mg",
|
||||
r"(\d)\s*[kK]b\b": r"\1 kb",
|
||||
r"(\d)\s*[mM]b\b": r"\1 Mb",
|
||||
r"(\d)\s*[gG]b\b": r"\1 Gb",
|
||||
r"(\d)\s*[pP]?[mM]\b": r"\1 pm",
|
||||
r"(\d)\s*[nN][mM]\b": r"\1 nm",
|
||||
r"(\d)\s*[cC][mM]\b": r"\1 cm",
|
||||
r"(\d)\s*[mM][lL]\b": r"\1 mL",
|
||||
r"(\d)\s*[uU][lL]\b": r"\1 µL",
|
||||
r"(\d)\s*°\s*[cC]\b": r"\1 °C",
|
||||
}
|
||||
for pattern, repl in unit_fixes.items():
|
||||
text = re.sub(pattern, repl, text)
|
||||
return text
|
||||
|
||||
|
||||
def fix_line_breaks(text: str) -> str:
|
||||
text = re.sub(r"(?<!\n)\n(?!\n)", " ", text)
|
||||
text = re.sub(r"\n{3,}", "\n\n", text)
|
||||
return text
|
||||
|
||||
|
||||
def clean_text(text: str) -> CleanResult:
|
||||
original_len = len(text)
|
||||
operations = []
|
||||
|
||||
text = fix_encoding(text)
|
||||
operations.append("fix_encoding")
|
||||
|
||||
text = normalize_unicode(text)
|
||||
operations.append("normalize_unicode")
|
||||
|
||||
text = fix_ligatures(text)
|
||||
operations.append("fix_ligatures")
|
||||
|
||||
text = remove_control_chars(text)
|
||||
operations.append("remove_control_chars")
|
||||
|
||||
text = fix_hyphenation(text)
|
||||
operations.append("fix_hyphenation")
|
||||
|
||||
text = normalize_whitespace(text)
|
||||
operations.append("normalize_whitespace")
|
||||
|
||||
text = remove_ocr_artifacts(text)
|
||||
operations.append("remove_ocr_artifacts")
|
||||
|
||||
text = fix_scientific_notation(text)
|
||||
operations.append("fix_scientific_notation")
|
||||
|
||||
text = normalize_units(text)
|
||||
operations.append("normalize_units")
|
||||
|
||||
text = fix_line_breaks(text)
|
||||
operations.append("fix_line_breaks")
|
||||
|
||||
cleaned_len = len(text)
|
||||
logger.debug("text_cleaned", original=original_len, cleaned=cleaned_len, ops=operations)
|
||||
|
||||
return CleanResult(
|
||||
text=text,
|
||||
original_length=original_len,
|
||||
cleaned_length=cleaned_len,
|
||||
operations=operations,
|
||||
)
|
||||
227
src/ocr_pipeline/postprocess/entities.py
Normal file
227
src/ocr_pipeline/postprocess/entities.py
Normal file
@@ -0,0 +1,227 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
import spacy
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Entity:
|
||||
text: str
|
||||
label: str
|
||||
start: int
|
||||
end: int
|
||||
confidence: float
|
||||
context: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class EntityResult:
|
||||
entities: list[Entity]
|
||||
text: str
|
||||
total_entities: int
|
||||
entities_by_type: dict[str, int]
|
||||
|
||||
|
||||
class ScientificEntityRecognizer:
|
||||
def __init__(self):
|
||||
self._nlp = None
|
||||
self._initialized = False
|
||||
|
||||
def _init(self):
|
||||
if self._initialized:
|
||||
return
|
||||
try:
|
||||
|
||||
self._nlp = spacy.load(settings.entities.model)
|
||||
self._nlp.max_length = 2_000_000
|
||||
self._initialized = True
|
||||
logger.info("scispacy_model_loaded", model=settings.entities.model)
|
||||
except Exception as e:
|
||||
logger.warning("scispacy_load_failed", error=str(e), model=settings.entities.model)
|
||||
self._fallback_init()
|
||||
self._initialized = True
|
||||
|
||||
def _fallback_init(self):
|
||||
try:
|
||||
self._nlp = spacy.blank("en")
|
||||
self._nlp.add_pipe("sentencizer")
|
||||
logger.info("using_fallback_blank_model")
|
||||
except Exception:
|
||||
self._nlp = None
|
||||
|
||||
def extract(self, text: str) -> EntityResult:
|
||||
if not settings.entities.enabled:
|
||||
return EntityResult(entities=[], text=text, total_entities=0, entities_by_type={})
|
||||
|
||||
self._init()
|
||||
if self._nlp is None:
|
||||
return EntityResult(entities=[], text=text, total_entities=0, entities_by_type={})
|
||||
|
||||
try:
|
||||
doc = self._nlp(text)
|
||||
entities = []
|
||||
|
||||
for ent in doc.ents:
|
||||
if ent.label_ in settings.entities.types:
|
||||
context_start = max(0, ent.start_char - 100)
|
||||
context_end = min(len(text), ent.end_char + 100)
|
||||
context = text[context_start:context_end].strip()
|
||||
|
||||
entities.append(
|
||||
Entity(
|
||||
text=ent.text,
|
||||
label=ent.label_,
|
||||
start=ent.start_char,
|
||||
end=ent.end_char,
|
||||
confidence=0.9,
|
||||
context=context,
|
||||
)
|
||||
)
|
||||
|
||||
entities = (
|
||||
self._merge_entities(entities, text) if settings.entities.merge_entities else entities
|
||||
)
|
||||
entities = self._deduplicate_entities(entities)
|
||||
|
||||
by_type = {}
|
||||
for e in entities:
|
||||
by_type[e.label] = by_type.get(e.label, 0) + 1
|
||||
|
||||
return EntityResult(
|
||||
entities=entities,
|
||||
text=text,
|
||||
total_entities=len(entities),
|
||||
entities_by_type=by_type,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("entity_extraction_failed", error=str(e))
|
||||
return EntityResult(entities=[], text=text, total_entities=0, entities_by_type={})
|
||||
|
||||
def _merge_entities(self, entities: list[Entity], text: str) -> list[Entity]:
|
||||
if not entities:
|
||||
return entities
|
||||
|
||||
entities.sort(key=lambda e: (e.start, -e.end))
|
||||
merged = [entities[0]]
|
||||
|
||||
for ent in entities[1:]:
|
||||
last = merged[-1]
|
||||
if ent.start <= last.end + 5 and ent.label == last.label:
|
||||
merged[-1] = Entity(
|
||||
text=last.text + text[last.end : ent.start] + ent.text,
|
||||
label=last.label,
|
||||
start=last.start,
|
||||
end=ent.end,
|
||||
confidence=max(last.confidence, ent.confidence),
|
||||
context=last.context,
|
||||
)
|
||||
else:
|
||||
merged.append(ent)
|
||||
|
||||
return merged
|
||||
|
||||
def _deduplicate_entities(self, entities: list[Entity]) -> list[Entity]:
|
||||
seen = set()
|
||||
unique = []
|
||||
for ent in entities:
|
||||
key = (ent.text.lower(), ent.label)
|
||||
if key not in seen:
|
||||
seen.add(key)
|
||||
unique.append(ent)
|
||||
return unique
|
||||
|
||||
|
||||
class RegexEntityRecognizer:
|
||||
PATTERNS = {
|
||||
"GENE": [
|
||||
r"\b(?:[A-Z]{2,5}\d*[A-Z]?|[A-Z]{3,10})\b",
|
||||
r"\b(?:BRCA[12]|TP53|EGFR|KRAS|MYC|PIK3CA|PTEN|APC|MLH1|MSH2|MSH6|PMS2|ATM|CHEK2|PALB2|RAD51|BARD1|BRIP1|CDH1|STK11|VHL|WT1|NF1|NF2|RB1|MEN1|RET|MET|ALK|ROS1|NTRK[123]|BRAF|NRAS|HRAS|IDH[12]|TERT|FGFR[1234]|KIT|PDGFRA|FLT3|JAK2|MPL|CALR|ASXL1|DNMT3A|TET2|SRSF2|U2AF1|SF3B1|ZRSR2|CBL|CSF3R|RUNX1|CEBPA|NPM1|FLT3|DNMT3A|IDH[12]|TET2|ASXL1|SRSF2|U2AF1|SF3B1|ZRSR2|CBL|CSF3R|RUNX1|CEBPA|NPM1)\b",
|
||||
],
|
||||
"PROTEIN": [
|
||||
r"\b(?:p53|p21|p16|CDK[1246]|cyclin\s+[ADEB]|Rb|E2F|MDM2|BAX|BCL[2XL]|CASP[389]|PARP|H2AX|ATM|ATR|CHK[12]|BRCA[12]|RAD51|FANCD2|FANCI|PCNA|RPA|MSH[26]|MLH1|PMS2|EXO1|MRE11|RAD50|NBN|CTIP|BRCA1|BRCA2|PALB2|RAD51C|RAD51D|BRIP1|BARD1)\b",
|
||||
],
|
||||
"CHEMICAL": [
|
||||
r"\b(?:DMSO|PBS|EDTA|Tris|HEPES|SDS|DTT|BME|NaCl|KCl|MgCl2|CaCl2|NaOH|HCl|H2SO4|HNO3|EtOH|MeOH|IPA|DMSO|DMF|DMA|THF|DCM|CHCl3|CH2Cl2|EtOAc|hexane|pentane|acetone|acetonitrile|water|H2O|buffer|media|serum|FBS|BSA|pen/strep|penicillin|streptomycin|trypsin|EDTA|collagenase|dispase|accutase)\b",
|
||||
r"\b(?:[A-Z][a-z]?\d*(?:\([^)]+\))?(?:\s*[+\-]\s*[A-Z][a-z]?\d*(?:\([^)]+\))?)*)\b",
|
||||
],
|
||||
"CONCENTRATION": [
|
||||
r"\b\d+(?:\.\d+)?\s*(?:[µumMpnc]?[Mm]|[µumMpnc]?[Mm]/[Ll]|[µumMpnc]?[Mm]\s*[Ll]?|[µumMpnc]?[gG]\s*/\s*[Ll]|[µumMpnc]?[gG]\s*/\s*[mM][lL]|[µumMpnc]?[gG]\s*/\s*[dL]|[µumMpnc]?[gG]\s*/\s*[mM][lL]|[µumMpnc]?[gG]\s*/\s*[dD][lL]|[µumMpnc]?[gG]\s*/\s*100\s*[mM][lL]|[µumMpnc]?[gG]\s*/\s*[kK][gG])\b",
|
||||
],
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def extract(cls, text: str) -> list[Entity]:
|
||||
entities = []
|
||||
for label, patterns in cls.PATTERNS.items():
|
||||
for pattern in patterns:
|
||||
for match in re.finditer(pattern, text, re.IGNORECASE):
|
||||
context_start = max(0, match.start() - 50)
|
||||
context_end = min(len(text), match.end() + 50)
|
||||
entities.append(
|
||||
Entity(
|
||||
text=match.group(),
|
||||
label=label,
|
||||
start=match.start(),
|
||||
end=match.end(),
|
||||
confidence=0.7,
|
||||
context=text[context_start:context_end].strip(),
|
||||
)
|
||||
)
|
||||
return entities
|
||||
|
||||
|
||||
def extract_entities(text: str) -> EntityResult:
|
||||
recognizer = ScientificEntityRecognizer()
|
||||
result = recognizer.extract(text)
|
||||
|
||||
regex_entities = RegexEntityRecognizer.extract(text)
|
||||
all_entities = result.entities + regex_entities
|
||||
all_entities = _deduplicate(all_entities)
|
||||
|
||||
by_type = {}
|
||||
for e in all_entities:
|
||||
by_type[e.label] = by_type.get(e.label, 0) + 1
|
||||
|
||||
return EntityResult(
|
||||
entities=all_entities,
|
||||
text=text,
|
||||
total_entities=len(all_entities),
|
||||
entities_by_type=by_type,
|
||||
)
|
||||
|
||||
|
||||
def _deduplicate(entities: list[Entity]) -> list[Entity]:
|
||||
seen = set()
|
||||
unique = []
|
||||
for ent in sorted(entities, key=lambda e: (e.start, -e.end)):
|
||||
key = (ent.text.lower(), ent.label)
|
||||
if key not in seen:
|
||||
seen.add(key)
|
||||
unique.append(ent)
|
||||
return unique
|
||||
|
||||
|
||||
def format_entities_for_markdown(result: EntityResult) -> str:
|
||||
if not result.entities:
|
||||
return "No entities detected."
|
||||
|
||||
lines = [
|
||||
"## Detected Entities",
|
||||
"",
|
||||
"| Entity | Type | Count | Contexts |",
|
||||
"|--------|------|-------|----------|",
|
||||
]
|
||||
|
||||
for label in sorted(result.entities_by_type.keys()):
|
||||
ents = [e for e in result.entities if e.label == label]
|
||||
contexts = {e.context[:80] for e in ents[:3]}
|
||||
lines.append(f"| {ents[0].text} | {label} | {len(ents)} | {'; '.join(contexts)} |")
|
||||
|
||||
return "\n".join(lines)
|
||||
18
src/ocr_pipeline/utils/__init__.py
Normal file
18
src/ocr_pipeline/utils/__init__.py
Normal file
@@ -0,0 +1,18 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from .db import Database, ProcessedFile, get_db
|
||||
from .logging import ProgressTracker, create_progress, get_logger, setup_logging
|
||||
from .migrate import compute_file_hash, migrate_notebook, parse_notebook_for_images
|
||||
|
||||
__all__ = [
|
||||
"setup_logging",
|
||||
"get_logger",
|
||||
"create_progress",
|
||||
"ProgressTracker",
|
||||
"parse_notebook_for_images",
|
||||
"migrate_notebook",
|
||||
"compute_file_hash",
|
||||
"get_db",
|
||||
"Database",
|
||||
"ProcessedFile",
|
||||
]
|
||||
BIN
src/ocr_pipeline/utils/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/utils/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/utils/__pycache__/db.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/utils/__pycache__/db.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/utils/__pycache__/logging.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/utils/__pycache__/logging.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/utils/__pycache__/migrate.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/utils/__pycache__/migrate.cpython-313.pyc
Normal file
Binary file not shown.
166
src/ocr_pipeline/utils/db.py
Normal file
166
src/ocr_pipeline/utils/db.py
Normal file
@@ -0,0 +1,166 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProcessedFile:
|
||||
path: str
|
||||
hash: str
|
||||
timestamp: float
|
||||
run_id: int
|
||||
size: int
|
||||
mtime: float
|
||||
|
||||
|
||||
class Database:
|
||||
def __init__(self, db_path: Path):
|
||||
self.db_path = Path(db_path).expanduser().resolve()
|
||||
self.db_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._init_db()
|
||||
|
||||
@contextmanager
|
||||
def _conn(self):
|
||||
conn = sqlite3.connect(self.db_path)
|
||||
conn.row_factory = sqlite3.Row
|
||||
try:
|
||||
yield conn
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
def _init_db(self):
|
||||
with self._conn() as conn:
|
||||
conn.executescript("""
|
||||
CREATE TABLE IF NOT EXISTS processed_files (
|
||||
hash TEXT PRIMARY KEY,
|
||||
path TEXT NOT NULL,
|
||||
timestamp TEXT NOT NULL,
|
||||
ocr_engine TEXT,
|
||||
confidence REAL,
|
||||
chunks_created INTEGER DEFAULT 0,
|
||||
created_at REAL DEFAULT (strftime('%s', 'now'))
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS runs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
date TEXT NOT NULL,
|
||||
run_number INTEGER NOT NULL,
|
||||
total_images INTEGER DEFAULT 0,
|
||||
successful INTEGER DEFAULT 0,
|
||||
failed INTEGER DEFAULT 0,
|
||||
chunks_created INTEGER DEFAULT 0,
|
||||
started_at REAL DEFAULT (strftime('%s', 'now')),
|
||||
completed_at REAL,
|
||||
UNIQUE(date, run_number)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_processed_hash ON processed_files(hash);
|
||||
CREATE INDEX IF NOT EXISTS idx_runs_date ON runs(date);
|
||||
""")
|
||||
|
||||
def is_processed(self, file_hash: str, path: Path | str | None = None) -> bool:
|
||||
"""Check if a file is already processed by hash (and optionally path)."""
|
||||
with self._conn() as conn:
|
||||
if path:
|
||||
cursor = conn.execute(
|
||||
"SELECT 1 FROM processed_files WHERE hash = ? AND path = ?",
|
||||
(file_hash, str(path)),
|
||||
)
|
||||
else:
|
||||
cursor = conn.execute(
|
||||
"SELECT 1 FROM processed_files WHERE hash = ?",
|
||||
(file_hash,),
|
||||
)
|
||||
return cursor.fetchone() is not None
|
||||
|
||||
def get_processed(self, file_hash: str) -> dict[str, Any] | None:
|
||||
with self._conn() as conn:
|
||||
cursor = conn.execute("SELECT * FROM processed_files WHERE hash = ?", (file_hash,))
|
||||
row = cursor.fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
def mark_processed(
|
||||
self,
|
||||
file_hash: str,
|
||||
path: str,
|
||||
timestamp: float | str,
|
||||
engine: str = "",
|
||||
confidence: float = 0.0,
|
||||
chunks: int = 0,
|
||||
run_id: int = 0,
|
||||
):
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"""INSERT OR REPLACE INTO processed_files
|
||||
(hash, path, timestamp, ocr_engine, confidence, chunks_created)
|
||||
VALUES (?, ?, ?, ?, ?, ?)""",
|
||||
(file_hash, path, str(timestamp), engine, confidence, chunks),
|
||||
)
|
||||
|
||||
def get_stats(self, days: int = 30) -> dict[str, Any]:
|
||||
with self._conn() as conn:
|
||||
cursor = conn.execute("""
|
||||
SELECT
|
||||
COUNT(*) as total_processed,
|
||||
SUM(chunks_created) as total_chunks,
|
||||
AVG(confidence) as avg_confidence
|
||||
FROM processed_files
|
||||
""")
|
||||
row = cursor.fetchone()
|
||||
return dict(row) if row else {}
|
||||
|
||||
def start_run(self, date: str | None = None) -> int:
|
||||
if date is None:
|
||||
date = time.strftime("%Y-%m-%d")
|
||||
with self._conn() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT COALESCE(MAX(run_number), 0) + 1 FROM runs WHERE date = ?", (date,)
|
||||
)
|
||||
run_number = cursor.fetchone()[0]
|
||||
conn.execute(
|
||||
"INSERT INTO runs (date, run_number, total_images) VALUES (?, ?, 0)",
|
||||
(date, run_number),
|
||||
)
|
||||
return run_number
|
||||
|
||||
def complete_run(self, date: str, run_number: int, stats: dict[str, int]):
|
||||
with self._conn() as conn:
|
||||
conn.execute(
|
||||
"""UPDATE runs SET
|
||||
total_images = ?,
|
||||
successful = ?,
|
||||
failed = ?,
|
||||
chunks_created = ?,
|
||||
completed_at = ?
|
||||
WHERE date = ? AND run_number = ?""",
|
||||
(
|
||||
stats.get("total", 0),
|
||||
stats.get("successful", 0),
|
||||
stats.get("failed", 0),
|
||||
stats.get("chunks", 0),
|
||||
time.time(),
|
||||
date,
|
||||
run_number,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
_db_instance: Database | None = None
|
||||
|
||||
|
||||
def get_db() -> Database:
|
||||
global _db_instance
|
||||
if _db_instance is None:
|
||||
from ocr_pipeline.config import settings
|
||||
|
||||
_db_instance = Database(settings.watch_db_path)
|
||||
return _db_instance
|
||||
132
src/ocr_pipeline/utils/logging.py
Normal file
132
src/ocr_pipeline/utils/logging.py
Normal file
@@ -0,0 +1,132 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
import structlog
|
||||
from rich.logging import RichHandler
|
||||
from rich.progress import (
|
||||
BarColumn,
|
||||
Progress,
|
||||
SpinnerColumn,
|
||||
TaskProgressColumn,
|
||||
TextColumn,
|
||||
TimeElapsedColumn,
|
||||
)
|
||||
|
||||
|
||||
def setup_logging(level: str = "INFO", log_file: Path | None = None) -> None:
|
||||
"""Configure structlog with rich console output and optional file logging."""
|
||||
log_level = getattr(logging, level.upper(), logging.INFO)
|
||||
|
||||
shared_processors = [
|
||||
structlog.contextvars.merge_contextvars,
|
||||
structlog.stdlib.add_logger_name,
|
||||
structlog.stdlib.add_log_level,
|
||||
structlog.processors.TimeStamper(fmt="ISO"),
|
||||
structlog.processors.StackInfoRenderer(),
|
||||
structlog.dev.set_exc_info,
|
||||
]
|
||||
|
||||
structlog.configure(
|
||||
processors=shared_processors
|
||||
+ [
|
||||
structlog.stdlib.ProcessorFormatter.wrap_for_formatter,
|
||||
],
|
||||
logger_factory=structlog.stdlib.LoggerFactory(),
|
||||
wrapper_class=structlog.stdlib.BoundLogger,
|
||||
context_class=dict,
|
||||
)
|
||||
|
||||
formatter = structlog.stdlib.ProcessorFormatter(
|
||||
foreign_pre_chain=shared_processors,
|
||||
processors=[
|
||||
structlog.stdlib.ProcessorFormatter.remove_processors_meta,
|
||||
structlog.dev.ConsoleRenderer(colors=True),
|
||||
],
|
||||
)
|
||||
|
||||
console_handler = RichHandler(
|
||||
rich_tracebacks=True,
|
||||
tracebacks_show_locals=False,
|
||||
markup=True,
|
||||
show_time=False,
|
||||
show_path=False,
|
||||
)
|
||||
console_handler.setFormatter(formatter)
|
||||
|
||||
handlers: list[logging.Handler] = [console_handler]
|
||||
|
||||
if log_file:
|
||||
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
file_handler = logging.FileHandler(log_file)
|
||||
file_formatter = structlog.stdlib.ProcessorFormatter(
|
||||
foreign_pre_chain=shared_processors,
|
||||
processors=[
|
||||
structlog.stdlib.ProcessorFormatter.remove_processors_meta,
|
||||
structlog.processors.JSONRenderer(),
|
||||
],
|
||||
)
|
||||
file_handler.setFormatter(file_formatter)
|
||||
handlers.append(file_handler)
|
||||
|
||||
root_logger = logging.getLogger()
|
||||
root_logger.handlers = []
|
||||
root_logger.setLevel(log_level)
|
||||
for handler in handlers:
|
||||
root_logger.addHandler(handler)
|
||||
|
||||
logging.getLogger("PIL").setLevel(logging.WARNING)
|
||||
logging.getLogger("paddle").setLevel(logging.WARNING)
|
||||
logging.getLogger("torch").setLevel(logging.WARNING)
|
||||
logging.getLogger("transformers").setLevel(logging.WARNING)
|
||||
logging.getLogger("watchdog").setLevel(logging.WARNING)
|
||||
|
||||
|
||||
def get_logger(name: str) -> structlog.BoundLogger:
|
||||
return structlog.get_logger(name)
|
||||
|
||||
|
||||
def create_progress() -> Progress:
|
||||
return Progress(
|
||||
SpinnerColumn(),
|
||||
TextColumn("[bold blue]{task.description}"),
|
||||
BarColumn(),
|
||||
TaskProgressColumn(),
|
||||
TimeElapsedColumn(),
|
||||
transient=True,
|
||||
)
|
||||
|
||||
|
||||
class ProgressTracker:
|
||||
def __init__(self, total: int, description: str = "Processing"):
|
||||
self.total = total
|
||||
self.description = description
|
||||
self.progress: Progress | None = None
|
||||
self.task_id: int | None = None
|
||||
self.completed = 0
|
||||
self.failed = 0
|
||||
self.logger = get_logger(__name__)
|
||||
|
||||
def __enter__(self) -> ProgressTracker:
|
||||
self.progress = create_progress()
|
||||
self.progress.start()
|
||||
self.task_id = self.progress.add_task(self.description, total=self.total)
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
|
||||
if self.progress:
|
||||
self.progress.stop()
|
||||
|
||||
def update(self, advance: int = 1, success: bool = True) -> None:
|
||||
if self.progress and self.task_id is not None:
|
||||
self.progress.update(self.task_id, advance=advance)
|
||||
self.completed += 1
|
||||
if not success:
|
||||
self.failed += 1
|
||||
|
||||
def log_result(self, path: Path, success: bool, error: str | None = None) -> None:
|
||||
if success:
|
||||
self.logger.info("processed", path=str(path))
|
||||
else:
|
||||
self.logger.error("failed", path=str(path), error=error)
|
||||
92
src/ocr_pipeline/utils/migrate.py
Normal file
92
src/ocr_pipeline/utils/migrate.py
Normal file
@@ -0,0 +1,92 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import nbformat
|
||||
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def parse_notebook_for_images(notebook_path: Path) -> list[Path]:
|
||||
image_paths = []
|
||||
try:
|
||||
with open(notebook_path) as f:
|
||||
nb = nbformat.read(f, as_version=4)
|
||||
|
||||
for cell in nb.cells:
|
||||
if cell.cell_type == "code":
|
||||
for output in cell.outputs:
|
||||
if output.output_type == "stream" and output.name == "stdout":
|
||||
text = output.text
|
||||
for line in text.split("\n"):
|
||||
line = line.strip()
|
||||
if line.startswith("/") and any(
|
||||
line.endswith(ext)
|
||||
for ext in (".png", ".jpg", ".jpeg", ".tiff", ".bmp")
|
||||
):
|
||||
path = Path(line)
|
||||
if path.exists():
|
||||
image_paths.append(path)
|
||||
elif "SCR-" in line and any(
|
||||
line.endswith(ext) for ext in (".png", ".jpg", ".jpeg")
|
||||
):
|
||||
for part in line.split():
|
||||
if part.startswith("SCR-") and any(
|
||||
part.endswith(ext) for ext in (".png", ".jpg", ".jpeg")
|
||||
):
|
||||
path = Path("/Users/Aman/Pictures") / part
|
||||
if path.exists():
|
||||
image_paths.append(path)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("notebook_parse_failed", path=str(notebook_path), error=str(e))
|
||||
|
||||
return list(set(image_paths))
|
||||
|
||||
|
||||
def compute_file_hash(path: Path) -> str:
|
||||
hasher = hashlib.sha256()
|
||||
with open(path, "rb") as f:
|
||||
for chunk in iter(lambda: f.read(8192), b""):
|
||||
hasher.update(chunk)
|
||||
return hasher.hexdigest()
|
||||
|
||||
|
||||
def migrate_notebook(notebook_path: Path, pipeline: Any) -> dict[str, Any]:
|
||||
logger.info("starting_notebook_migration", notebook=str(notebook_path))
|
||||
|
||||
image_paths = parse_notebook_for_images(notebook_path)
|
||||
logger.info("found_images_in_notebook", count=len(image_paths))
|
||||
|
||||
results = {
|
||||
"total": len(image_paths),
|
||||
"processed": 0,
|
||||
"skipped": 0,
|
||||
"failed": 0,
|
||||
"errors": [],
|
||||
}
|
||||
|
||||
for img_path in image_paths:
|
||||
try:
|
||||
img_hash = compute_file_hash(img_path)
|
||||
existing = pipeline.db.get_processed(img_hash)
|
||||
if existing:
|
||||
logger.info("skipping_already_processed", path=str(img_path), hash=img_hash[:16])
|
||||
results["skipped"] += 1
|
||||
continue
|
||||
|
||||
logger.info("migrating_image", path=str(img_path))
|
||||
pipeline.process_single(img_path)
|
||||
results["processed"] += 1
|
||||
|
||||
except Exception as e:
|
||||
logger.error("migration_failed", path=str(img_path), error=str(e))
|
||||
results["failed"] += 1
|
||||
results["errors"].append({"path": str(img_path), "error": str(e)})
|
||||
|
||||
logger.info("migration_complete", **results)
|
||||
return results
|
||||
144
src/ocr_pipeline/watch.py
Normal file
144
src/ocr_pipeline/watch.py
Normal file
@@ -0,0 +1,144 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
from watchdog.events import FileSystemEventHandler
|
||||
from watchdog.observers import Observer
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.db import ProcessedFile, get_db
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class ScreenshotHandler(FileSystemEventHandler):
|
||||
def __init__(self, process_callback: Callable[[Path], None]):
|
||||
self.process_callback = process_callback
|
||||
self.pending: dict[str, float] = {}
|
||||
self.lock = threading.Lock()
|
||||
self.debounce = settings.watch.debounce_seconds
|
||||
|
||||
def _should_process(self, path: Path) -> bool:
|
||||
if not path.is_file():
|
||||
return False
|
||||
if path.suffix.lower() not in (".png", ".jpg", ".jpeg", ".tiff", ".bmp"):
|
||||
return False
|
||||
for pattern in settings.watch.ignore_patterns:
|
||||
if path.match(pattern):
|
||||
return False
|
||||
return True
|
||||
|
||||
def on_created(self, event):
|
||||
if not event.is_directory:
|
||||
path = Path(event.src_path)
|
||||
if self._should_process(path):
|
||||
with self.lock:
|
||||
self.pending[str(path)] = time.time()
|
||||
|
||||
def on_modified(self, event):
|
||||
if not event.is_directory:
|
||||
path = Path(event.src_path)
|
||||
if self._should_process(path):
|
||||
with self.lock:
|
||||
self.pending[str(path)] = time.time()
|
||||
|
||||
def check_pending(self) -> list[Path]:
|
||||
now = time.time()
|
||||
to_process = []
|
||||
|
||||
with self.lock:
|
||||
for path_str, timestamp in list(self.pending.items()):
|
||||
if now - timestamp >= self.debounce:
|
||||
path = Path(path_str)
|
||||
if path.exists():
|
||||
to_process.append(path)
|
||||
del self.pending[path_str]
|
||||
|
||||
return to_process
|
||||
|
||||
|
||||
class Watcher:
|
||||
def __init__(self, process_callback: Callable[[Path], None]):
|
||||
self.process_callback = process_callback
|
||||
self.observer: Observer | None = None
|
||||
self.handler: ScreenshotHandler | None = None
|
||||
self.running = False
|
||||
self.db = get_db()
|
||||
self.run_id: int | None = None
|
||||
|
||||
def start(self) -> None:
|
||||
if self.running:
|
||||
return
|
||||
|
||||
self.run_id = self.db.start_run()
|
||||
self.running = True
|
||||
|
||||
self.handler = ScreenshotHandler(self._process_file)
|
||||
self.observer = Observer()
|
||||
|
||||
for base_path_str in settings.input.paths:
|
||||
base_path = Path(base_path_str).expanduser()
|
||||
if base_path.exists():
|
||||
self.observer.schedule(
|
||||
self.handler, str(base_path), recursive=settings.input.recursive
|
||||
)
|
||||
logger.info("watching_directory", path=str(base_path))
|
||||
|
||||
self.observer.start()
|
||||
|
||||
self._watch_thread = threading.Thread(target=self._watch_loop, daemon=True)
|
||||
self._watch_thread.start()
|
||||
|
||||
logger.info("watcher_started", run_id=self.run_id)
|
||||
|
||||
def _watch_loop(self) -> None:
|
||||
poll_interval = settings.watch.poll_interval
|
||||
processed = 0
|
||||
failed = 0
|
||||
|
||||
while self.running:
|
||||
time.sleep(poll_interval)
|
||||
|
||||
if self.handler:
|
||||
for path in self.handler.check_pending():
|
||||
try:
|
||||
self._process_file(path)
|
||||
processed += 1
|
||||
except Exception as e:
|
||||
logger.error("watch_process_failed", path=str(path), error=str(e))
|
||||
failed += 1
|
||||
|
||||
self.db.end_run(self.run_id, processed, failed)
|
||||
logger.info("watcher_stopped", run_id=self.run_id, processed=processed, failed=failed)
|
||||
|
||||
def _process_file(self, path: Path) -> None:
|
||||
from ocr_pipeline.ocr.engine import compute_image_hash
|
||||
|
||||
file_hash = compute_image_hash(path)
|
||||
if self.db.is_processed(path, file_hash):
|
||||
logger.debug("file_already_processed", path=str(path))
|
||||
return
|
||||
|
||||
logger.info("processing_new_file", path=str(path))
|
||||
self.process_callback(path)
|
||||
self.db.mark_processed(
|
||||
ProcessedFile(
|
||||
path=str(path),
|
||||
hash=file_hash,
|
||||
timestamp=time.time(),
|
||||
run_id=self.run_id,
|
||||
size=path.stat().st_size,
|
||||
mtime=path.stat().st_mtime,
|
||||
)
|
||||
)
|
||||
|
||||
def stop(self) -> None:
|
||||
self.running = False
|
||||
if self.observer:
|
||||
self.observer.stop()
|
||||
self.observer.join(timeout=5)
|
||||
logger.info("watcher_stopped")
|
||||
5
src/ocr_pipeline/watch/__init__.py
Normal file
5
src/ocr_pipeline/watch/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from .watcher import Watcher
|
||||
|
||||
__all__ = ["Watcher"]
|
||||
BIN
src/ocr_pipeline/watch/__pycache__/__init__.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/watch/__pycache__/__init__.cpython-313.pyc
Normal file
Binary file not shown.
BIN
src/ocr_pipeline/watch/__pycache__/watcher.cpython-313.pyc
Normal file
BIN
src/ocr_pipeline/watch/__pycache__/watcher.cpython-313.pyc
Normal file
Binary file not shown.
144
src/ocr_pipeline/watch/watcher.py
Normal file
144
src/ocr_pipeline/watch/watcher.py
Normal file
@@ -0,0 +1,144 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
from watchdog.events import FileSystemEventHandler
|
||||
from watchdog.observers import Observer
|
||||
|
||||
from ocr_pipeline.config import settings
|
||||
from ocr_pipeline.utils.db import get_db
|
||||
from ocr_pipeline.utils.logging import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class ScreenshotHandler(FileSystemEventHandler):
|
||||
def __init__(self, process_callback: Callable[[Path], None]):
|
||||
self.process_callback = process_callback
|
||||
self.pending: dict[str, float] = {}
|
||||
self.lock = threading.Lock()
|
||||
self.debounce = settings.watch.debounce_seconds
|
||||
|
||||
def _should_process(self, path: Path) -> bool:
|
||||
if not path.is_file():
|
||||
return False
|
||||
if path.suffix.lower() not in (".png", ".jpg", ".jpeg", ".tiff", ".bmp"):
|
||||
return False
|
||||
for pattern in settings.watch.ignore_patterns:
|
||||
if path.match(pattern):
|
||||
return False
|
||||
return True
|
||||
|
||||
def on_created(self, event):
|
||||
if not event.is_directory:
|
||||
path = Path(event.src_path)
|
||||
if self._should_process(path):
|
||||
with self.lock:
|
||||
self.pending[str(path)] = time.time()
|
||||
|
||||
def on_modified(self, event):
|
||||
if not event.is_directory:
|
||||
path = Path(event.src_path)
|
||||
if self._should_process(path):
|
||||
with self.lock:
|
||||
self.pending[str(path)] = time.time()
|
||||
|
||||
def check_pending(self) -> list[Path]:
|
||||
now = time.time()
|
||||
to_process = []
|
||||
|
||||
with self.lock:
|
||||
for path_str, timestamp in list(self.pending.items()):
|
||||
if now - timestamp >= self.debounce:
|
||||
path = Path(path_str)
|
||||
if path.exists():
|
||||
to_process.append(path)
|
||||
del self.pending[path_str]
|
||||
|
||||
return to_process
|
||||
|
||||
|
||||
class Watcher:
|
||||
def __init__(self, process_callback: Callable[[Path], None]):
|
||||
self.process_callback = process_callback
|
||||
self.observer: Observer | None = None
|
||||
self.handler: ScreenshotHandler | None = None
|
||||
self.running = False
|
||||
self.db = get_db()
|
||||
self.run_id: int | None = None
|
||||
|
||||
def start(self) -> None:
|
||||
if self.running:
|
||||
return
|
||||
|
||||
self.run_id = self.db.start_run()
|
||||
self.running = True
|
||||
|
||||
self.handler = ScreenshotHandler(self._process_file)
|
||||
self.observer = Observer()
|
||||
|
||||
for base_path_str in settings.input.paths:
|
||||
base_path = Path(base_path_str).expanduser()
|
||||
if base_path.exists():
|
||||
self.observer.schedule(
|
||||
self.handler, str(base_path), recursive=settings.input.recursive
|
||||
)
|
||||
logger.info("watching_directory", path=str(base_path))
|
||||
|
||||
self.observer.start()
|
||||
|
||||
self._watch_thread = threading.Thread(target=self._watch_loop, daemon=True)
|
||||
self._watch_thread.start()
|
||||
|
||||
logger.info("watcher_started", run_id=self.run_id)
|
||||
|
||||
def _watch_loop(self) -> None:
|
||||
poll_interval = settings.watch.poll_interval
|
||||
processed = 0
|
||||
failed = 0
|
||||
|
||||
while self.running:
|
||||
time.sleep(poll_interval)
|
||||
|
||||
if self.handler:
|
||||
for path in self.handler.check_pending():
|
||||
try:
|
||||
self._process_file(path)
|
||||
processed += 1
|
||||
except Exception as e:
|
||||
logger.error("watch_process_failed", path=str(path), error=str(e))
|
||||
failed += 1
|
||||
|
||||
self.db.complete_run(
|
||||
time.strftime("%Y-%m-%d"),
|
||||
self.run_id,
|
||||
{"total": processed, "successful": processed, "failed": failed, "chunks": 0},
|
||||
)
|
||||
logger.info("watcher_stopped", run_id=self.run_id, processed=processed, failed=failed)
|
||||
|
||||
def _process_file(self, path: Path) -> None:
|
||||
from ocr_pipeline.ocr.engine import compute_image_hash
|
||||
|
||||
file_hash = compute_image_hash(path)
|
||||
if self.db.is_processed(file_hash, path):
|
||||
logger.debug("file_already_processed", path=str(path))
|
||||
return
|
||||
|
||||
logger.info("processing_new_file", path=str(path))
|
||||
self.process_callback(path)
|
||||
self.db.mark_processed(
|
||||
file_hash=file_hash,
|
||||
path=str(path),
|
||||
timestamp=time.time(),
|
||||
run_id=self.run_id,
|
||||
)
|
||||
|
||||
def stop(self) -> None:
|
||||
self.running = False
|
||||
if self.observer:
|
||||
self.observer.stop()
|
||||
self.observer.join(timeout=5)
|
||||
logger.info("watcher_stop_requested")
|
||||
0
tests/__init__.py
Normal file
0
tests/__init__.py
Normal file
4
tests/conftest.py
Normal file
4
tests/conftest.py
Normal file
@@ -0,0 +1,4 @@
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
|
||||
73
tests/test_postprocess.py
Normal file
73
tests/test_postprocess.py
Normal file
@@ -0,0 +1,73 @@
|
||||
import pytest
|
||||
from pathlib import Path
|
||||
|
||||
from ocr_pipeline.postprocess.clean import clean_text, CleanResult
|
||||
from ocr_pipeline.postprocess.chunk import chunk_text, TextChunk
|
||||
from ocr_pipeline.config import settings
|
||||
|
||||
|
||||
class TestTextCleaning:
|
||||
def test_fix_encoding(self):
|
||||
result = clean_text("café")
|
||||
assert "café" in result.text or "cafe" in result.text
|
||||
|
||||
def test_normalize_whitespace(self):
|
||||
result = clean_text("hello world\n\n\n\nfoo")
|
||||
assert "hello world" in result.text
|
||||
assert "foo" in result.text
|
||||
|
||||
def test_fix_hyphenation(self):
|
||||
result = clean_text("hyphen-\nated")
|
||||
assert "hyphenated" in result.text
|
||||
|
||||
def test_scientific_notation(self):
|
||||
result = clean_text("1.5 x 10 ^ -3")
|
||||
assert "×10^-3" in result.text or "x10^-3" in result.text
|
||||
|
||||
def test_normalize_units(self):
|
||||
result = clean_text("50 uM and 100 ug/mL")
|
||||
assert "µM" in result.text
|
||||
assert "µg/mL" in result.text or "µg/mL" in result.text
|
||||
|
||||
def test_remove_ocr_artifacts(self):
|
||||
result = clean_text("normal text\n|||\n───\nmore text")
|
||||
assert "normal text" in result.text
|
||||
assert "more text" in result.text
|
||||
|
||||
|
||||
class TestChunking:
|
||||
def test_basic_chunking(self):
|
||||
text = "A" * 500 + "\n\n" + "B" * 500 + "\n\n" + "C" * 500
|
||||
chunks = chunk_text(text)
|
||||
assert len(chunks) >= 2
|
||||
assert all(isinstance(c, TextChunk) for c in chunks)
|
||||
|
||||
def test_empty_text(self):
|
||||
chunks = chunk_text("")
|
||||
assert chunks == []
|
||||
|
||||
def test_single_chunk(self):
|
||||
text = "Short text that fits in one chunk"
|
||||
chunks = chunk_text(text)
|
||||
assert len(chunks) == 1
|
||||
|
||||
def test_metadata_propagation(self):
|
||||
text = "A" * 1500
|
||||
meta = {"source": "test.png", "custom": "value"}
|
||||
chunks = chunk_text(text, meta)
|
||||
assert all(c.metadata.get("source") == "test.png" for c in chunks)
|
||||
assert all(c.metadata.get("custom") == "value" for c in chunks)
|
||||
|
||||
|
||||
class TestConfig:
|
||||
def test_settings_load(self):
|
||||
assert settings.input.paths
|
||||
assert settings.ocr.engine in ("paddleocr", "tesseract", "auto")
|
||||
assert settings.output.base_directory
|
||||
|
||||
def test_output_dir_property(self):
|
||||
assert settings.output_dir.exists() or settings.output_dir.parent.exists()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
Reference in New Issue
Block a user