SPaCial_server_api / correction_post_processor.py
cooldan's picture
Upload 5 files
df339ab verified
Raw History Blame Contribute Delete
5.95 kB
#!/usr/bin/env python3
"""
Correction Post-Processor for OCR Results
Handles text corrections and improvements based on rules
"""
import json
import re
import logging
from pathlib import Path
logger = logging.getLogger(__name__)
class CorrectionPostProcessor:
"""Post-process OCR results with correction rules"""
def __init__(self, rules_file='correction_rules.json'):
self.rules_file = rules_file
self.rules = self.load_rules()
self.stats = {
'total_processed': 0,
'total_corrected': 0,
'text_fixed': 0,
'box_moved': 0,
'new_zone': 0,
'deleted': 0,
'validated': 0
}
def load_rules(self):
"""Load correction rules from JSON file"""
try:
if Path(self.rules_file).exists():
with open(self.rules_file, 'r', encoding='utf-8') as f:
rules = json.load(f)
logger.info(f"Loaded {len(rules.get('text_replacements', {}))} correction rules")
return rules
else:
logger.warning(f"Rules file {self.rules_file} not found, using default rules")
return self.get_default_rules()
except Exception as e:
logger.error(f"Failed to load correction rules: {e}")
return self.get_default_rules()
def get_default_rules(self):
"""Get default correction rules"""
return {
"text_replacements": {
# Common OCR mistakes
"O": "0", # Letter O to number 0 in dimensions
"I": "1", # Letter I to number 1 in dimensions
"l": "1", # lowercase l to number 1
"S": "5", # Letter S to number 5
"G": "6", # Letter G to number 6
"B": "8", # Letter B to number 8
# Dimension-specific corrections
"Ø": "Ø", # Keep diameter symbol
"±": "±", # Keep plus-minus symbol
# Remove common OCR artifacts
"|": "", # Remove stray vertical bars
"~": "", # Remove tildes
"`": "", # Remove backticks
},
"dimension_patterns": [
r"(\d+\.?\d*)\s*±\s*(\d+\.?\d*)", # ± tolerance
r"(\d+\.?\d*)\s*\+\s*(\d+\.?\d*)\s*/\s*-\s*(\d+\.?\d*)", # +tolerance/-tolerance
r"(\d+\.?\d*)\s*-\s*(\d+\.?\d*)", # -tolerance only
r"(\d+\.?\d*)\s*\+\s*(\d+\.?\d*)", # +tolerance only
],
"confidence_threshold": 0.7,
"min_text_length": 1
}
def process_zones(self, zones):
"""Process OCR zones with correction rules"""
if not zones:
return zones
self.stats['total_processed'] = len(zones)
corrected_zones = []
for zone in zones:
try:
corrected_zone = self.correct_zone(zone)
if corrected_zone:
corrected_zones.append(corrected_zone)
# Track corrections
if self.was_corrected(zone, corrected_zone):
self.stats['total_corrected'] += 1
except Exception as e:
logger.warning(f"Failed to correct zone: {e}")
corrected_zones.append(zone)
logger.info(f"Processed {self.stats['total_processed']} zones, corrected {self.stats['total_corrected']}")
return corrected_zones
def correct_zone(self, zone):
"""Apply corrections to a single zone"""
if not zone or not zone.get('text'):
return zone
original_text = zone['text']
corrected_text = self.correct_text(original_text)
# Create corrected zone
corrected_zone = zone.copy()
corrected_zone['text'] = corrected_text
corrected_zone['original_text'] = original_text
corrected_zone['was_corrected'] = corrected_text != original_text
return corrected_zone
def correct_text(self, text):
"""Apply text corrections based on rules"""
if not text:
return text
corrected = text
# Apply character replacements
replacements = self.rules.get('text_replacements', {})
for wrong, correct in replacements.items():
corrected = corrected.replace(wrong, correct)
# Apply regex patterns for dimensions
dimension_patterns = self.rules.get('dimension_patterns', [])
for pattern in dimension_patterns:
# This is just for validation, actual parsing is done elsewhere
if re.search(pattern, corrected):
break
# Clean up whitespace
corrected = re.sub(r'\s+', ' ', corrected).strip()
return corrected
def was_corrected(self, original_zone, corrected_zone):
"""Check if zone was actually corrected"""
if not original_zone or not corrected_zone:
return False
original_text = original_zone.get('text', '')
corrected_text = corrected_zone.get('text', '')
return original_text != corrected_text
def get_stats(self):
"""Get correction statistics"""
return self.stats.copy()
def reset_stats(self):
"""Reset statistics"""
self.stats = {
'total_processed': 0,
'total_corrected': 0,
'text_fixed': 0,
'box_moved': 0,
'new_zone': 0,
'deleted': 0,
'validated': 0
}