Team Ai
Apppublic

documentExtractionag051/ExtractDocument

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
form_field_extractor.py733 linesDownload Raw Back to root
1# form_field_extractor.py2"""3Form Field Extractor Module4 5Extracts key-value form fields from OCR data produced by pdf_utils.py.6Supports multiple matching strategies:7- Regex patterns8- Line-based matching  9- Key-value proximity detection10 11Usage:12    extractor = FormFieldExtractor()13    results = extractor.extract_form_fields(14        form_fields=["Invoice Number", "Invoice Date", "Total Amount"],15        ocr_data=ocr_blocks16    )17"""18 19import re20import logging21from typing import List, Dict, Any, Optional, Tuple22from dataclasses import dataclass, field23 24# Set up logging25logging.basicConfig(level=logging.INFO, format='%(levelname)s:%(name)s:%(message)s')26logger = logging.getLogger(__name__)27 28 29# =============================================================================30# Data Classes31# =============================================================================32 33@dataclass34class ExtractedField:35    """Represents an extracted form field with its value and metadata."""36    field_name: str37    value: str38    confidence: float = 1.039    geometry: Optional[Dict[str, int]] = None40    page_num: int = 141    match_type: str = "unknown"  # regex, inline, proximity, pattern42    43    def to_dict(self) -> Dict[str, Any]:44        """Convert to dictionary format."""45        result = {46            "field_name": self.field_name,47            "value": self.value,48            "confidence": self.confidence,49            "page_num": self.page_num,50            "match_type": self.match_type51        }52        if self.geometry:53            result["geometry"] = self.geometry54            result["bounds"] = f"{self.geometry['x1']}, {self.geometry['y1']}, " \55                              f"{self.geometry['x2'] - self.geometry['x1']}, " \56                              f"{self.geometry['y2'] - self.geometry['y1']}"57        return result58 59 60# =============================================================================61# Field Variation Patterns62# =============================================================================63 64# Common variations for field names (key -> list of variations)65FIELD_VARIATIONS = {66    "invoice number": ["invoice no", "invoice #", "inv no", "inv #", "invoice num", "inv number"],67    "invoice date": ["inv date", "date of invoice", "invoice dt"],68    "po number": ["po no", "po #", "purchase order", "purchase order no", "p.o. number"],69    "order number": ["order no", "order #", "order id"],70    "customer name": ["customer", "cust name", "bill to", "sold to"],71    "vendor name": ["vendor", "supplier", "supplier name"],72    "total amount": ["total", "grand total", "net total", "amount due", "total due"],73    "subtotal": ["sub total", "sub-total", "net amount"],74    "tax": ["tax amount", "gst", "vat", "sales tax", "tax total"],75    "gstin": ["gst no", "gst number", "gstin no", "gst in"],76    "due date": ["payment due", "due by", "pay by"],77    "ship to": ["shipping address", "deliver to", "ship address"],78    "bill to": ["billing address", "invoice to", "billed to"],79}80 81# Separators between key and value82KEY_VALUE_SEPARATORS = [':', '-', 'โ€“', '|', '=', '.']83 84 85# =============================================================================86# FormFieldExtractor Class87# =============================================================================88 89class FormFieldExtractor:90    """91    Extracts form fields (key-value pairs) from OCR data.92    93    Supports multiple extraction strategies:94    1. Inline extraction: "Invoice Number: 12345" in same line95    2. Proximity extraction: Key and value in separate but adjacent elements96    3. Pattern-based extraction: Using regex patterns for specific value types97    """98    99    def __init__(100        self,101        max_horizontal_gap: float = 0.15,  # 15% of page width for proximity102        max_vertical_gap: float = 0.03,     # 3% of page height for proximity103        case_sensitive: bool = False104    ):105        """106        Initialize the form field extractor.107        108        Args:109            max_horizontal_gap: Maximum horizontal gap (% of page width) for value detection110            max_vertical_gap: Maximum vertical gap (% of page height) for value detection111            case_sensitive: Whether field matching should be case-sensitive112        """113        self.max_horizontal_gap = max_horizontal_gap114        self.max_vertical_gap = max_vertical_gap115        self.case_sensitive = case_sensitive116        117        # Build variation lookup118        self._variation_map = self._build_variation_map()119    120    def _build_variation_map(self) -> Dict[str, str]:121        """Build a map from variation -> canonical field name."""122        variation_map = {}123        for canonical, variations in FIELD_VARIATIONS.items():124            variation_map[canonical] = canonical125            for var in variations:126                variation_map[var] = canonical127        return variation_map128    129    def extract_form_fields(130        self,131        form_fields: List[str],132        ocr_data: List[Dict[str, Any]],133        page_dimensions: Optional[Dict[int, Tuple[int, int]]] = None134    ) -> Dict[str, Any]:135        """136        Extract and return key-value pairs for requested form fields.137        138        Args:139            form_fields: List of field names to extract (e.g., ["Invoice Number", "Invoice Date"])140            ocr_data: List of OCR blocks from pdf_utils.py output141            page_dimensions: Optional dict of {page_num: (width, height)}142            143        Returns:144            Dict with extracted form fields:145            {146                "Invoice Number": {"value": "INV-2024-001", "confidence": 0.95, ...},147                "Invoice Date": {"value": "2024-01-15", "confidence": 0.90, ...},148                ...149            }150        """151        logger.info(f"Extracting {len(form_fields)} form fields from {len(ocr_data)} OCR blocks")152        153        # Filter to LINE elements only (more reliable than WORD)154        line_blocks = [155            block for block in ocr_data 156            if block.get('blockType') == 'LINE'157        ]158        logger.info(f"Found {len(line_blocks)} LINE blocks")159        160        # Infer page dimensions if not provided161        if not page_dimensions:162            page_dimensions = self._infer_page_dimensions(ocr_data)163        164        # Extract each requested field165        results = {}166        for field_name in form_fields:167            extracted = self._extract_single_field(field_name, line_blocks, page_dimensions)168            if extracted:169                results[field_name] = extracted.to_dict()170            else:171                # Return empty value if not found172                results[field_name] = {173                    "field_name": field_name,174                    "value": "",175                    "confidence": 0.0,176                    "match_type": "not_found"177                }178        179        logger.info(f"Extracted {sum(1 for v in results.values() if v['value'])} of {len(form_fields)} fields")180        return results181    182    def _extract_single_field(183        self,184        field_name: str,185        line_blocks: List[Dict[str, Any]],186        page_dimensions: Dict[int, Tuple[int, int]]187    ) -> Optional[ExtractedField]:188        """189        Extract a single form field using multiple strategies.190        191        Strategies (in order of priority):192        1. Inline extraction with separator193        2. Pattern-based extraction (for known patterns)194        3. Proximity-based extraction195        """196        # Normalize field name for matching197        field_norm = self._normalize_text(field_name)198        199        # Get all variations to search for200        variations = self._get_field_variations(field_name)201        logger.debug(f"Searching for '{field_name}' with variations: {variations}")202        203        # Strategy 1: Inline extraction (key: value in same line)204        result = self._extract_inline(field_name, variations, line_blocks)205        if result:206            logger.debug(f"Found '{field_name}' via inline extraction: {result.value}")207            return result208        209        # Strategy 2: Pattern-based extraction210        result = self._extract_by_pattern(field_name, variations, line_blocks)211        if result:212            logger.debug(f"Found '{field_name}' via pattern extraction: {result.value}")213            return result214        215        # Strategy 3: Proximity-based extraction216        result = self._extract_by_proximity(field_name, variations, line_blocks, page_dimensions)217        if result:218            logger.debug(f"Found '{field_name}' via proximity extraction: {result.value}")219            return result220        221        logger.debug(f"Could not find value for '{field_name}'")222        return None223    224    def _extract_inline(225        self,226        field_name: str,227        variations: List[str],228        line_blocks: List[Dict[str, Any]]229    ) -> Optional[ExtractedField]:230        """231        Extract field where key and value are in the same line.232        E.g., "Invoice Number: INV-2024-001" or "Invoice Number - INV-2024-001"233        """234        for block in line_blocks:235            text = block.get('text', '')236            text_norm = self._normalize_text(text)237            238            for variation in variations:239                var_norm = self._normalize_text(variation)240                241                # Check if line contains the field name242                if var_norm not in text_norm:243                    continue244                245                # Try to extract value after separator246                for sep in KEY_VALUE_SEPARATORS:247                    # Build pattern: field_name + separator + value248                    pattern = re.compile(249                        rf'{re.escape(var_norm)}\s*{re.escape(sep)}\s*(.+)',250                        re.IGNORECASE251                    )252                    match = pattern.search(text)253                    if match:254                        value = match.group(1).strip()255                        # Clean up value (remove trailing separators)256                        value = self._clean_value(value)257                        if value:258                            return ExtractedField(259                                field_name=field_name,260                                value=value,261                                confidence=0.95,262                                geometry=block.get('geometry'),263                                page_num=block.get('pageNum', 1),264                                match_type="inline"265                            )266                267                # Try without separator (field name followed by value)268                # E.g., "Invoice Number INV-2024-001"269                pattern = re.compile(270                    rf'{re.escape(var_norm)}\s+(\S+.*?)$',271                    re.IGNORECASE272                )273                match = pattern.search(text)274                if match:275                    value = match.group(1).strip()276                    value = self._clean_value(value)277                    # Only accept if value doesn't look like another field name278                    if value and not self._looks_like_field_name(value):279                        return ExtractedField(280                            field_name=field_name,281                            value=value,282                            confidence=0.85,283                            geometry=block.get('geometry'),284                            page_num=block.get('pageNum', 1),285                            match_type="inline_no_sep"286                        )287        288        return None289    290    def _extract_by_pattern(291        self,292        field_name: str,293        variations: List[str],294        line_blocks: List[Dict[str, Any]]295    ) -> Optional[ExtractedField]:296        """297        Extract field using known value patterns.298        E.g., dates, amounts, reference numbers.299        """300        field_lower = field_name.lower()301        302        # Define patterns for specific field types303        patterns = {}304        305        # Date patterns306        if any(kw in field_lower for kw in ['date', 'dt']):307            patterns['date'] = [308                r'\d{1,2}[/-]\d{1,2}[/-]\d{2,4}',309                r'\d{4}[/-]\d{1,2}[/-]\d{1,2}',310                r'(?:Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec)[a-z]*[\s.,]+\d{1,2}[\s.,]+\d{2,4}',311                r'\d{1,2}[\s]+(?:Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec)[a-z]*[\s]+\d{2,4}',312            ]313        314        # Amount patterns315        if any(kw in field_lower for kw in ['amount', 'total', 'price', 'value', 'tax', 'subtotal']):316            patterns['amount'] = [317                r'[\$โ‚ฌยฃโ‚น][\d,]+\.?\d*',318                r'[\d,]+\.?\d*\s*(?:USD|EUR|GBP|INR)',319                r'(?:Rs\.?|INR)\s*[\d,]+\.?\d*',320            ]321        322        # Reference number patterns323        if any(kw in field_lower for kw in ['number', 'no', '#', 'id', 'ref']):324            patterns['reference'] = [325                r'[A-Z]{2,4}[-/]?\d{4,}[-/]?\d*',  # INV-2024-001, PO/12345326                r'\d{6,}',  # Pure numeric references327            ]328        329        # GSTIN pattern330        if 'gstin' in field_lower or 'gst' in field_lower:331            patterns['gstin'] = [332                r'\d{2}[A-Z]{5}\d{4}[A-Z]{1}[A-Z\d]{1}[Z]{1}[A-Z\d]{1}',  # Indian GSTIN333            ]334        335        if not patterns:336            return None337        338        # Search for field name first, then look for pattern nearby339        for block in line_blocks:340            text = block.get('text', '')341            text_norm = self._normalize_text(text)342            343            for variation in variations:344                var_norm = self._normalize_text(variation)345                346                if var_norm in text_norm:347                    # Found field name, look for pattern in same line348                    for pattern_type, pattern_list in patterns.items():349                        for pattern in pattern_list:350                            # Look for pattern after the field name351                            idx = text_norm.find(var_norm)352                            search_text = text[idx + len(variation):]353                            354                            match = re.search(pattern, search_text, re.IGNORECASE)355                            if match:356                                value = match.group(0).strip()357                                return ExtractedField(358                                    field_name=field_name,359                                    value=value,360                                    confidence=0.90,361                                    geometry=block.get('geometry'),362                                    page_num=block.get('pageNum', 1),363                                    match_type=f"pattern_{pattern_type}"364                                )365        366        return None367    368    def _extract_by_proximity(369        self,370        field_name: str,371        variations: List[str],372        line_blocks: List[Dict[str, Any]],373        page_dimensions: Dict[int, Tuple[int, int]]374    ) -> Optional[ExtractedField]:375        """376        Extract field where key and value are in separate but adjacent elements.377        """378        # Find blocks containing the field name379        label_blocks = []380        for block in line_blocks:381            text = block.get('text', '')382            text_norm = self._normalize_text(text)383            384            for variation in variations:385                var_norm = self._normalize_text(variation)386                if var_norm in text_norm or text_norm in var_norm:387                    # Check if this looks like a label (ends with separator or is the field name)388                    if self._looks_like_label(text, variation):389                        label_blocks.append((block, variation))390                        break391        392        if not label_blocks:393            return None394        395        # For each label, find the best value candidate396        for label_block, matched_variation in label_blocks:397            label_geom = label_block.get('geometry', {})398            label_page = label_block.get('pageNum', 1)399            400            page_dims = page_dimensions.get(label_page, (2479, 3508))401            page_width, page_height = page_dims402            403            max_h_gap = page_width * self.max_horizontal_gap404            max_v_gap = page_height * self.max_vertical_gap405            406            best_value = None407            best_score = 0408            409            for block in line_blocks:410                if block.get('pageNum', 1) != label_page:411                    continue412                if block.get('text', '').strip() == label_block.get('text', '').strip():413                    continue414                415                value_geom = block.get('geometry', {})416                value_text = block.get('text', '').strip()417                418                # Skip if looks like another label419                if self._looks_like_field_name(value_text):420                    continue421                422                # Calculate position score423                score = self._calculate_proximity_score(424                    label_geom, value_geom, max_h_gap, max_v_gap425                )426                427                if score > best_score:428                    best_score = score429                    best_value = block430            431            if best_value and best_score > 0.3:432                value_text = best_value.get('text', '').strip()433                value_text = self._clean_value(value_text)434                435                return ExtractedField(436                    field_name=field_name,437                    value=value_text,438                    confidence=round(best_score, 2),439                    geometry=best_value.get('geometry'),440                    page_num=best_value.get('pageNum', 1),441                    match_type="proximity"442                )443        444        return None445    446    def _calculate_proximity_score(447        self,448        label_geom: Dict[str, int],449        value_geom: Dict[str, int],450        max_h_gap: float,451        max_v_gap: float452    ) -> float:453        """Calculate how likely a value is associated with a label based on position."""454        if not label_geom or not value_geom:455            return 0.0456        457        label_x1 = label_geom.get('x1', 0)458        label_x2 = label_geom.get('x2', 0)459        label_y1 = label_geom.get('y1', 0)460        label_y2 = label_geom.get('y2', 0)461        462        value_x1 = value_geom.get('x1', 0)463        value_x2 = value_geom.get('x2', 0)464        value_y1 = value_geom.get('y1', 0)465        value_y2 = value_geom.get('y2', 0)466        467        score = 0.0468        469        # Check if value is to the right of label (same row)470        label_center_y = (label_y1 + label_y2) / 2471        value_center_y = (value_y1 + value_y2) / 2472        473        vertical_overlap = (min(label_y2, value_y2) - max(label_y1, value_y1))474        min_height = min(label_y2 - label_y1, value_y2 - value_y1)475        476        if min_height > 0 and vertical_overlap / min_height > 0.5:477            # Same row - value should be to the right478            if value_x1 >= label_x2:479                h_gap = value_x1 - label_x2480                if h_gap < max_h_gap:481                    proximity = 1.0 - (h_gap / max_h_gap)482                    score = max(score, 0.5 + 0.4 * proximity)483        484        # Check if value is below label485        if value_y1 >= label_y2:486            v_gap = value_y1 - label_y2487            if v_gap < max_v_gap:488                # Check horizontal alignment (value should be near label x position)489                horizontal_overlap = (min(label_x2, value_x2) - max(label_x1, value_x1))490                if horizontal_overlap > 0 or abs(value_x1 - label_x1) < max_h_gap:491                    proximity = 1.0 - (v_gap / max_v_gap)492                    score = max(score, 0.4 + 0.3 * proximity)493        494        return score495    496    def _get_field_variations(self, field_name: str) -> List[str]:497        """Get all variations of a field name to search for."""498        field_norm = self._normalize_text(field_name)499        variations = [field_name]500        501        # Check if this maps to a canonical name502        canonical = self._variation_map.get(field_norm)503        if canonical:504            # Add all known variations505            variations.extend(FIELD_VARIATIONS.get(canonical, []))506        507        # Add common transformations508        variations.append(field_name.replace(' ', ''))  # NoSpaces509        variations.append(field_name.replace(' ', '_'))  # With_underscores510        511        # Deduplicate while preserving order512        seen = set()513        unique = []514        for v in variations:515            v_norm = self._normalize_text(v)516            if v_norm not in seen:517                seen.add(v_norm)518                unique.append(v)519        520        return unique521    522    def _normalize_text(self, text: str) -> str:523        """Normalize text for comparison."""524        if not text:525            return ""526        result = ' '.join(text.lower().strip().split())527        return result528    529    def _clean_value(self, value: str) -> str:530        """Clean extracted value."""531        if not value:532            return ""533        534        # Remove leading/trailing separators535        value = value.strip()536        for sep in KEY_VALUE_SEPARATORS + [',', ';']:537            value = value.strip(sep).strip()538        539        return value540    541    def _looks_like_label(self, text: str, field_name: str) -> bool:542        """Check if text looks like a field label."""543        text = text.strip()544        545        # Ends with separator546        for sep in KEY_VALUE_SEPARATORS:547            if text.endswith(sep):548                return True549        550        # Is exactly the field name (possibly with separator)551        text_clean = text.rstrip(':- ')552        if self._normalize_text(text_clean) == self._normalize_text(field_name):553            return True554        555        return True  # Default to treating as label if field name matches556    557    def _looks_like_field_name(self, text: str) -> bool:558        """Check if text looks like a field label (not a value)."""559        text = text.strip()560        561        # Very short text might be a value562        if len(text) < 3:563            return False564        565        # Contains separator at end566        for sep in KEY_VALUE_SEPARATORS:567            if text.endswith(sep):568                return True569        570        # Is a known field variation571        text_norm = self._normalize_text(text.rstrip(':- '))572        if text_norm in self._variation_map:573            return True574        575        return False576    577    def _infer_page_dimensions(self, ocr_data: List[Dict[str, Any]]) -> Dict[int, Tuple[int, int]]:578        """Infer page dimensions from OCR data."""579        page_dims = {}580        581        for block in ocr_data:582            page_num = block.get('pageNum', 1)583            geom = block.get('geometry', {})584            585            if page_num not in page_dims:586                page_dims[page_num] = (0, 0)587            588            current_w, current_h = page_dims[page_num]589            page_dims[page_num] = (590                max(current_w, geom.get('x2', 0)),591                max(current_h, geom.get('y2', 0))592            )593        594        # Add buffer for page dimensions595        for page_num in page_dims:596            w, h = page_dims[page_num]597            page_dims[page_num] = (int(w * 1.1), int(h * 1.1))598        599        return page_dims if page_dims else {1: (2479, 3508)}600 601 602# =============================================================================603# Standalone Functions604# =============================================================================605 606def extract_form_fields(607    form_fields: List[str],608    ocr_data: List[Dict[str, Any]],609    page_dimensions: Optional[Dict[int, Tuple[int, int]]] = None610) -> Dict[str, Any]:611    """612    Extract and return key-value pairs for requested form fields.613    614    This is a convenience function that creates a FormFieldExtractor instance615    and calls its extract_form_fields method.616    617    Args:618        form_fields: List of field names to extract (e.g., ["Invoice Number", "Invoice Date"])619        ocr_data: List of OCR blocks from pdf_utils.py output (ocrBlocks array)620        page_dimensions: Optional dict of {page_num: (width, height)}621        622    Returns:623        Dict with extracted form fields624    """625    extractor = FormFieldExtractor()626    return extractor.extract_form_fields(form_fields, ocr_data, page_dimensions)627 628 629def merge_extraction_results(630    pdf_output: Dict[str, Any],631    form_fields_result: Dict[str, Any],632    table_fields_result: Optional[Dict[str, Any]] = None633) -> Dict[str, Any]:634    """635    Merge form field and table extraction results with PDF output.636    637    Args:638        pdf_output: Original output from pdf_utils.py639        form_fields_result: Result from extract_form_fields()640        table_fields_result: Optional result from table extraction641        642    Returns:643        Merged output with structure:644        {645            "version": "1.0",646            "metadata": {...},647            "pages": [...],648            "ocrBlocks": {649                "formFields": {...},650                "tableFields": {...}651            }652        }653    """654    result = {655        "version": pdf_output.get("version", "1.0"),656        "metadata": pdf_output.get("metadata", {}),657        "pages": pdf_output.get("pages", []),658        "ocrBlocks": {659            "formFields": form_fields_result,660            "tableFields": table_fields_result or {}661        }662    }663    664    return result665 666 667# =============================================================================668# CLI Entry Point669# =============================================================================670 671if __name__ == '__main__':672    import sys673    import json674    675    # Example usage676    print("Form Field Extractor")677    print("=" * 50)678    679    if len(sys.argv) < 2:680        print("Usage: python form_field_extractor.py <json_file> [field1,field2,...]")681        print("\nExample:")682        print("  python form_field_extractor.py output.json \"Invoice Number,Invoice Date,Total Amount\"")683        sys.exit(1)684    685    json_file = sys.argv[1]686    687    # Default fields to extract688    default_fields = ["Invoice Number", "Invoice Date", "Total Amount", "GSTIN", "PO Number"]689    690    if len(sys.argv) > 2:691        fields_str = sys.argv[2]692        form_fields = [f.strip() for f in fields_str.split(',')]693    else:694        form_fields = default_fields695    696    print(f"Input file: {json_file}")697    print(f"Fields to extract: {form_fields}")698    print("-" * 50)699    700    # Load JSON701    with open(json_file, 'r') as f:702        data = json.load(f)703    704    # Extract form fields705    ocr_blocks = data.get('ocrBlocks', [])706    707    # Get page dimensions708    page_dimensions = {}709    for page in data.get('pages', []):710        page_num = page.get('pageNum', len(page_dimensions) + 1)711        page_dimensions[page_num] = (page.get('width', 2479), page.get('height', 3508))712    713    # Run extraction714    extractor = FormFieldExtractor()715    results = extractor.extract_form_fields(form_fields, ocr_blocks, page_dimensions)716    717    # Print results718    print("\nExtracted Fields:")719    print("-" * 50)720    for field_name, field_data in results.items():721        value = field_data.get('value', '')722        match_type = field_data.get('match_type', 'unknown')723        confidence = field_data.get('confidence', 0)724        725        status = "โœ“" if value else "โœ—"726        print(f"{status} {field_name}: {value or '(not found)'}")727        if value:728            print(f"   Match: {match_type}, Confidence: {confidence}")729    730    print("\n" + "-" * 50)731    print("JSON Output:")732    print(json.dumps(results, indent=2))733