Document Question Answering
Transformers
PyTorch
English
document-processing
ocr
ner
text-classification
information-extraction
invoice
receipt
form
Instructions to use mrrobot2610/IDP-Machine-learning with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use mrrobot2610/IDP-Machine-learning with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("document-question-answering", model="mrrobot2610/IDP-Machine-learning")# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("mrrobot2610/IDP-Machine-learning", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download postprocessing.py from mrrobot2610/IDP-Machine-learning: direct link, hf CLI and curl.
- Browser
- Download file 14.2 kB
-
https://huggingface.co/mrrobot2610/IDP-Machine-learning/resolve/main/postprocessing.py
- Command line
-
hf download hf://mrrobot2610/IDP-Machine-learning/postprocessing.py
-
curl -L -o postprocessing.py https://huggingface.co/mrrobot2610/IDP-Machine-learning/resolve/main/postprocessing.py
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)) | |