IDP-Machine-learning / postprocessing.py
mrrobot2610's picture
Initial commit: IDP (Intelligent Document Processing) System
1a7ee60
Raw History Blame Contribute Delete
14.2 kB
"""
Post-Processing and Key-Value Extraction
Combines NER predictions with regex patterns and heuristics
Implements confidence scoring and field validation
"""
import re
from typing import Dict, List, Optional, Tuple
import logging
from datetime import datetime
import numpy as np
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class PostProcessor:
"""
Post-processing layer for IDP
Combines NER outputs with regex patterns and validation rules
"""
def __init__(self):
# Regex patterns for common fields
self.patterns = {
'invoice_number': [
r'INV[-\s]?\d+',
r'Invoice\s*#?\s*:?\s*([A-Z0-9\-]+)',
r'\b\d{6,}\b', # 6+ digits
],
'date': [
r'\d{1,2}[/-]\d{1,2}[/-]\d{2,4}',
r'\d{4}[/-]\d{1,2}[/-]\d{1,2}',
r'\d{1,2}\s+(?:Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec)[a-z]*\s+\d{2,4}',
],
'amount': [
r'[₹$€£¥]\s*[\d,]+\.?\d*',
r'[\d,]+\.?\d*\s*[₹$€£¥]',
r'\b\d{1,3}(?:,\d{3})*(?:\.\d{2})?\b',
],
'gst_id': [
r'\d{2}[A-Z]{5}\d{4}[A-Z]{1}[A-Z\d]{1}[Z]{1}[A-Z\d]{1}', # Indian GST
r'[A-Z]{2}\d{9}', # VAT pattern
],
'email': [
r'\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,}\b',
],
'phone': [
r'\+?\d{1,3}[-.\s]?\(?\d{1,4}\)?[-.\s]?\d{1,4}[-.\s]?\d{1,9}',
],
}
def process_document(
self,
document_type: str,
classification_confidence: float,
ocr_text: str,
ocr_boxes: List[Dict],
ner_entities: List[Dict]
) -> Dict:
"""
Main post-processing pipeline
Args:
document_type: Predicted document type (INVOICE, RECEIPT, FORM, OTHER)
classification_confidence: Confidence score from classifier
ocr_text: Full OCR extracted text
ocr_boxes: OCR bounding boxes with confidence
ner_entities: NER predicted entities
Returns:
Structured document with extracted fields
"""
logger.info(f"Post-processing {document_type} document")
# Extract fields based on document type
if document_type == 'INVOICE':
fields = self._extract_invoice_fields(ocr_text, ocr_boxes, ner_entities)
elif document_type == 'RECEIPT':
fields = self._extract_receipt_fields(ocr_text, ocr_boxes, ner_entities)
elif document_type == 'FORM':
fields = self._extract_form_fields(ocr_text, ocr_boxes, ner_entities)
elif document_type == 'BANK_STATEMENT':
fields = self._extract_bank_statement_fields(ocr_text, ocr_boxes, ner_entities)
else:
fields = self._extract_generic_fields(ocr_text, ocr_boxes, ner_entities)
# Build structured output
result = {
'document_type': document_type,
'classification_confidence': classification_confidence,
'fields': fields,
'raw_ocr_text': ocr_text[:1000], # First 1000 chars for debugging
}
return result
def _extract_invoice_fields(
self,
text: str,
boxes: List[Dict],
entities: List[Dict]
) -> Dict:
"""Extract invoice-specific fields"""
fields = {}
# Invoice number
invoice_num = self._extract_field(
'INVOICE_NUMBER',
text,
boxes,
entities,
self.patterns['invoice_number']
)
if invoice_num:
fields['invoice_number'] = invoice_num
# Date
date = self._extract_field(
'DATE',
text,
boxes,
entities,
self.patterns['date']
)
if date:
fields['date'] = self._normalize_date(date)
# Total amount
total = self._extract_field(
'TOTAL_AMOUNT',
text,
boxes,
entities,
self.patterns['amount']
)
if total:
fields['total_amount'] = self._normalize_amount(total)
# Tax amount
tax = self._extract_field(
'TAX_AMOUNT',
text,
boxes,
entities,
self.patterns['amount']
)
if tax:
fields['tax_amount'] = self._normalize_amount(tax)
# Vendor name
vendor = self._extract_field(
'VENDOR_NAME',
text,
boxes,
entities,
[]
)
if vendor:
fields['vendor_name'] = vendor
# Customer name
customer = self._extract_field(
'CUSTOMER_NAME',
text,
boxes,
entities,
[]
)
if customer:
fields['customer_name'] = customer
# GST/Tax ID
gst = self._extract_field(
'GST_ID',
text,
boxes,
entities,
self.patterns['gst_id']
)
if gst:
fields['gst_id'] = gst
return fields
def _extract_receipt_fields(
self,
text: str,
boxes: List[Dict],
entities: List[Dict]
) -> Dict:
"""Extract receipt-specific fields"""
# Similar to invoice but may have different priorities
return self._extract_invoice_fields(text, boxes, entities)
def _extract_form_fields(
self,
text: str,
boxes: List[Dict],
entities: List[Dict]
) -> Dict:
"""Extract form-specific fields"""
fields = {}
# Generic entity extraction
for entity in entities:
entity_type = entity['entity'].lower()
if entity_type not in fields:
fields[entity_type] = {
'value': entity['text'],
'confidence': entity['confidence'],
'bbox': [entity['start'], 0, entity['end'], 0] # Simplified
}
return fields
def _extract_bank_statement_fields(
self,
text: str,
boxes: List[Dict],
entities: List[Dict]
) -> Dict:
"""Extract bank statement specific fields"""
# Reuse invoice extraction logic for common fields like Date and Amount
return self._extract_invoice_fields(text, boxes, entities)
def _extract_generic_fields(
self,
text: str,
boxes: List[Dict],
entities: List[Dict]
) -> Dict:
"""Extract generic fields for unknown document types"""
return self._extract_form_fields(text, boxes, entities)
def _extract_field(
self,
entity_type: str,
text: str,
boxes: List[Dict],
entities: List[Dict],
fallback_patterns: List[str]
) -> Optional[Dict]:
"""
Extract a single field using NER + regex + heuristics
Returns:
Dictionary with value, confidence, and bbox, or None
"""
# Step 1: Check NER entities
for entity in entities:
if entity['entity'] == entity_type:
# Get OCR confidence for this text span
ocr_conf = self._get_ocr_confidence(entity, boxes)
# Combined confidence
combined_conf = (entity['confidence'] + ocr_conf) / 2
return {
'value': entity['text'],
'confidence': float(combined_conf),
'bbox': self._get_bbox_for_entity(entity, boxes),
'source': 'ner'
}
# Step 2: Fallback to regex patterns
if fallback_patterns:
for pattern in fallback_patterns:
matches = re.finditer(pattern, text, re.IGNORECASE)
for match in matches:
matched_text = match.group(0)
# Validate match
if self._validate_field(entity_type, matched_text):
# Find in OCR boxes
ocr_conf = self._find_in_boxes(matched_text, boxes)
return {
'value': matched_text,
'confidence': float(ocr_conf * 0.8), # Lower confidence for regex
'bbox': None,
'source': 'regex'
}
return None
def _validate_field(self, entity_type: str, value: str) -> bool:
"""Validate extracted field value"""
if not value or not value.strip():
return False
# Type-specific validation
if entity_type == 'DATE':
return len(value) >= 6 # Minimum date length
elif entity_type in ['TOTAL_AMOUNT', 'TAX_AMOUNT']:
return any(c.isdigit() for c in value)
elif entity_type == 'INVOICE_NUMBER':
return len(value) >= 3
return True
def _get_ocr_confidence(self, entity: Dict, boxes: List[Dict]) -> float:
"""Get average OCR confidence for entity text"""
entity_text = entity['text'].lower()
confidences = []
for box in boxes:
if box['text'].lower() in entity_text or entity_text in box['text'].lower():
confidences.append(box['confidence'])
return float(np.mean(confidences)) if confidences else 0.5
def _get_bbox_for_entity(self, entity: Dict, boxes: List[Dict]) -> Optional[List[float]]:
"""Get bounding box for entity from OCR boxes"""
entity_text = entity['text'].lower()
for box in boxes:
if entity_text in box['text'].lower():
return box['bbox']
return None
def _find_in_boxes(self, text: str, boxes: List[Dict]) -> float:
"""Find text in OCR boxes and return confidence"""
text_lower = text.lower()
for box in boxes:
if text_lower in box['text'].lower():
return float(box['confidence'])
return 0.5 # Default confidence
def _normalize_date(self, date_field: Dict) -> Dict:
"""Normalize date to standard format"""
date_str = date_field['value']
# Try to parse date
formats = [
'%d/%m/%Y', '%d-%m-%Y', '%Y-%m-%d', '%Y/%m/%d',
'%d/%m/%y', '%d-%m-%y',
'%d %B %Y', '%d %b %Y',
]
for fmt in formats:
try:
parsed = datetime.strptime(date_str, fmt)
date_field['value'] = parsed.strftime('%Y-%m-%d')
date_field['normalized'] = True
return date_field
except:
continue
# Could not parse, return as-is
date_field['normalized'] = False
return date_field
def _normalize_amount(self, amount_field: Dict) -> Dict:
"""Normalize amount to numeric value"""
amount_str = amount_field['value']
# Extract currency symbol
currency_symbols = {'₹': 'INR', '$': 'USD', '€': 'EUR', '£': 'GBP', '¥': 'JPY'}
currency = None
for symbol, code in currency_symbols.items():
if symbol in amount_str:
currency = code
amount_str = amount_str.replace(symbol, '')
break
# Remove commas and spaces
amount_str = amount_str.replace(',', '').replace(' ', '').strip()
# Extract numeric value
try:
numeric_value = float(amount_str)
amount_field['value'] = f"{numeric_value:.2f}"
amount_field['numeric_value'] = numeric_value
if currency:
amount_field['currency'] = currency
amount_field['normalized'] = True
except:
amount_field['normalized'] = False
return amount_field
if __name__ == "__main__":
# Example usage
processor = PostProcessor()
# Dummy data
ocr_text = """
INVOICE
Invoice #: INV-12345
Date: 28/11/2025
Total Amount: ₹ 12,500.00
Tax: ₹ 2,250.00
GST ID: 29ABCDE1234F1Z5
"""
ocr_boxes = [
{'text': 'INVOICE', 'bbox': [100, 50, 200, 80], 'confidence': 0.98},
{'text': 'INV-12345', 'bbox': [150, 100, 250, 120], 'confidence': 0.95},
{'text': '28/11/2025', 'bbox': [150, 130, 250, 150], 'confidence': 0.92},
{'text': '₹ 12,500.00', 'bbox': [200, 200, 300, 220], 'confidence': 0.96},
]
ner_entities = [
{'entity': 'INVOICE_NUMBER', 'text': 'INV-12345', 'confidence': 0.94, 'start': 30, 'end': 39},
{'entity': 'DATE', 'text': '28/11/2025', 'confidence': 0.89, 'start': 47, 'end': 57},
{'entity': 'TOTAL_AMOUNT', 'text': '₹ 12,500.00', 'confidence': 0.93, 'start': 75, 'end': 86},
]
result = processor.process_document(
document_type='INVOICE',
classification_confidence=0.96,
ocr_text=ocr_text,
ocr_boxes=ocr_boxes,
ner_entities=ner_entities
)
print("\nPost-processing result:")
print(json.dumps(result, indent=2))