""" Dataset Loading and Preprocessing for IDP Training Supports CORD-v2, SROIE, and FUNSD datasets from Local Archives and Hugging Face Converts to formats suitable for classification and NER training """ import io import json import logging import os import glob from typing import Dict, List, Optional, Tuple, Union import numpy as np from datasets import load_dataset, Dataset from PIL import Image logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) class CORDDatasetLoader: """Load and process CORD-v2 dataset from local archive""" def __init__(self, base_path: str): self.base_path = base_path def load_dataset_splits(self): """Load CORD-v2 dataset from local files""" logger.info(f"Loading CORD-v2 dataset from {self.base_path}...") splits = {} for split_name in ["train", "dev", "test"]: # Map 'val' to 'dev' if needed, but directory is 'dev' dir_name = "dev" if split_name == "validation" else split_name if split_name == "val": dir_name = "dev" json_dir = os.path.join(self.base_path, dir_name, "json") if not os.path.exists(json_dir): logger.warning(f"Split directory not found: {json_dir}") continue data = [] try: files = os.listdir(json_dir) json_files = [os.path.join(json_dir, f) for f in files if f.endswith('.json')] except Exception as e: logger.warning(f"Error listing dir {json_dir}: {e}") json_files = [] for json_file in json_files: try: with open(json_file, 'r', encoding='utf-8') as f: content = json.load(f) # Add filename as id content['id'] = os.path.basename(json_file).replace('.json', '') data.append(content) except Exception as e: logger.warning(f"Error reading {json_file}: {e}") splits[split_name] = data logger.info(f"Loaded {len(data)} examples for split '{split_name}'") return splits @staticmethod def extract_classification_data(dataset_split) -> List[Dict]: """ Extract data for document classification """ classification_data = [] for example in dataset_split: try: # CORD local JSON structure text_lines = [] if "valid_line" in example: for line in example["valid_line"]: for word in line.get("words", []): if "text" in word: text_lines.append(word["text"]) text = " ".join(text_lines) classification_data.append( { "text": text, "label": "RECEIPT", "image": None, # Image loading not implemented for local yet "doc_id": example.get("id", ""), } ) except Exception as e: doc_id = example.get("id", "unknown") logger.warning( f"Skipping record {doc_id} in CORD due to parsing error: {e}" ) continue return classification_data @staticmethod def extract_ner_data(dataset_split) -> List[Dict]: """ Extract data for NER training from local CORD JSON """ ner_data = [] # Mapping CORD fields to our entity types entity_mapping = { "menu.nm": "VENDOR_NAME", # Sometimes menu name is used as vendor/item "menu.nm": "O", # Actually menu items are not usually vendor names in general receipt NER, but let's keep consistent with previous logic if possible. # Previous logic: "menu.nm": "VENDOR_NAME". Wait, menu.nm is usually the item name. # Let's map strictly important fields. "total.total_price": "TOTAL_AMOUNT", "total.tax_price": "TAX_AMOUNT", "sub_total.subtotal_price": "TOTAL_AMOUNT", # "menu.price": "TOTAL_AMOUNT", # Individual prices are not total } # Refined mapping based on CORD categories # CORD categories: menu.nm, menu.cnt, menu.price, sub_total.subtotal_price, total.total_price, etc. for example in dataset_split: if "valid_line" not in example: continue tokens = [] labels = [] for line in example["valid_line"]: category = line.get("category", "O") # Map category to our label # We need to be careful. CORD has hierarchical categories. label_type = "O" if category in entity_mapping: label_type = entity_mapping[category] elif category == "menu.nm": # In the previous code it was VENDOR_NAME, but that seems wrong for menu items. # However, to maintain compatibility with the 'ner_labels' defined in UnifiedDatasetLoader, # we should map to what we have. # Available: INVOICE_NUMBER, DATE, TOTAL_AMOUNT, TAX_AMOUNT, VENDOR_NAME, CUSTOMER_NAME, ADDRESS, GST_ID # CORD is mostly food receipts. # Let's try to find mappings. pass # Check for other fields manually if not in mapping if category == "total.total_price": label_type = "TOTAL_AMOUNT" elif category == "total.tax_price": label_type = "TAX_AMOUNT" # For now, let's stick to a simple mapping or "O" if not sure, to avoid noise. for word in line.get("words", []): text = word.get("text", "") if not text: continue # Simple tokenization by space if needed, but usually 'text' is a word word_tokens = text.split() tokens.extend(word_tokens) if label_type != "O": labels.append(f"B-{label_type}") labels.extend([f"I-{label_type}"] * (len(word_tokens) - 1)) else: labels.extend(["O"] * len(word_tokens)) if tokens: ner_data.append( { "tokens": tokens, "labels": labels, "doc_id": example.get("id", ""), } ) return ner_data class SROIEDatasetLoader: """Load and process SROIE dataset from local archive""" def __init__(self, base_path: str): self.base_path = base_path def load_dataset_splits(self): """Load SROIE dataset from local files""" logger.info(f"Loading SROIE dataset from {self.base_path}...") splits = {} # SROIE structure: train/entities, train/box, train/img # We will treat 'train' as train and maybe split later, or look for 'test' folder for split_name in ["train", "test"]: split_dir = os.path.join(self.base_path, split_name) if not os.path.exists(split_dir): continue entities_dir = os.path.join(split_dir, "entities") box_dir = os.path.join(split_dir, "box") data = [] try: files = os.listdir(entities_dir) entity_files = [os.path.join(entities_dir, f) for f in files if f.endswith('.txt')] except Exception as e: logger.warning(f"Error listing dir {entities_dir}: {e}") entity_files = [] for entity_file in entity_files: file_id = os.path.basename(entity_file).replace('.txt', '') box_file = os.path.join(box_dir, f"{file_id}.txt") if not os.path.exists(box_file): continue try: # Read entities (Ground Truth) with open(entity_file, 'r', encoding='utf-8') as f: # SROIE entities are usually one line JSON entities = json.load(f) # Read boxes and text words = [] with open(box_file, 'r', encoding='utf-8') as f: for line in f: parts = line.strip().split(',') if len(parts) >= 9: # x1,y1,x2,y2,x3,y3,x4,y4,text # text might contain commas, so join the rest text = ",".join(parts[8:]) words.append(text) data.append({ "id": file_id, "entities": entities, "text_lines": words, "full_text": " ".join(words) }) except Exception as e: logger.warning(f"Error reading SROIE file {file_id}: {e}") splits[split_name] = data logger.info(f"Loaded {len(data)} examples for split '{split_name}'") return splits @staticmethod def extract_classification_data(dataset_split) -> List[Dict]: """Extract classification data (all RECEIPT)""" classification_data = [] for example in dataset_split: classification_data.append( { "text": example.get("full_text", ""), "label": "RECEIPT", "image": None, "doc_id": example.get("id", ""), } ) return classification_data @staticmethod def extract_ner_data(dataset_split) -> List[Dict]: """Extract NER data from SROIE""" ner_data = [] # SROIE keys: company, date, address, total key_mapping = { "company": "VENDOR_NAME", "date": "DATE", "address": "ADDRESS", "total": "TOTAL_AMOUNT" } for example in dataset_split: text_lines = example.get("text_lines", []) entities = example.get("entities", {}) # This is a hard problem: mapping loose entities to tokens in the text. # For SROIE, the 'entities' file gives the *value* of the field. # We need to find that value in the 'text_lines'. tokens = [] labels = [] # Flatten text lines into tokens all_tokens = [] for line in text_lines: all_tokens.extend(line.split()) # Initialize all labels to O token_labels = ["O"] * len(all_tokens) # Try to match entities # This is a naive matching approach full_text_tokens = all_tokens for key, value in entities.items(): if key not in key_mapping: continue target_label = key_mapping[key] value_tokens = value.split() if not value_tokens: continue # Find sequence of value_tokens in full_text_tokens len_val = len(value_tokens) for i in range(len(full_text_tokens) - len_val + 1): # Check match (case insensitive? SROIE is usually exact match but OCR might vary) # Let's try exact match first match = True for j in range(len_val): if full_text_tokens[i+j] != value_tokens[j]: match = False break if match: token_labels[i] = f"B-{target_label}" for k in range(1, len_val): token_labels[i+k] = f"I-{target_label}" # We only match the first occurrence for now break ner_data.append({ "tokens": full_text_tokens, "labels": token_labels, "doc_id": example.get("id", "") }) return ner_data class FUNSDDatasetLoader: """Load and process FUNSD dataset (forms) - Keeping Hugging Face for now as not in local archive""" @staticmethod def load_dataset_splits(): """Load FUNSD dataset from Hugging Face""" logger.info("Loading FUNSD dataset...") ds = load_dataset("nielsr/funsd") return ds @staticmethod def extract_classification_data(dataset_split) -> List[Dict]: """Extract classification data (all FORM)""" classification_data = [] for example in dataset_split: words = example.get("words", []) text = " ".join(words) if words else "" classification_data.append( { "text": text, "label": "FORM", "image": example.get("image"), "doc_id": example.get("id", ""), } ) return classification_data @staticmethod def extract_ner_data(dataset_split) -> List[Dict]: """Extract NER data from FUNSD""" ner_data = [] ner_tags_map = { 0: "O", 1: "B-HEADER", 2: "I-HEADER", 3: "B-QUESTION", 4: "I-QUESTION", 5: "B-ANSWER", 6: "I-ANSWER", } for example in dataset_split: words = example.get("words", []) ner_tags = example.get("ner_tags", []) if not words: continue labels = [ner_tags_map.get(tag, "O") for tag in ner_tags] ner_data.append( { "tokens": words, "labels": labels, "doc_id": example.get("id", ""), } ) return ner_data class UnifiedDatasetLoader: """Unified interface for loading all datasets""" def __init__(self): # Define local paths self.cord_path = r"c:/Users/Harsh-Stu/IDP[ML]/archive (2)/CORD" self.sroie_path = r"c:/Users/Harsh-Stu/IDP[ML]/archive (1)/SROIE2019" self.loaders = { "cord": CORDDatasetLoader(self.cord_path), "sroie": SROIEDatasetLoader(self.sroie_path), "funsd": FUNSDDatasetLoader(), } def load_classification_dataset( self, datasets: List[str] = ["cord"], split: str = "train" ) -> List[Dict]: """ Load and combine classification data from multiple datasets """ all_data = [] for dataset_name in datasets: if dataset_name not in self.loaders: logger.warning(f"Unknown dataset: {dataset_name}") continue try: loader = self.loaders[dataset_name] ds = loader.load_dataset_splits() # Handle different split names target_split = split if split == "validation": if "val" in ds: target_split = "val" elif "dev" in ds: target_split = "dev" elif "test" in ds: target_split = "test" # Fallback if target_split not in ds and split == "train": # Fallback for train if not exact match (unlikely) pass if target_split in ds: data = loader.extract_classification_data(ds[target_split]) all_data.extend(data) logger.info( f"Loaded {len(data)} examples from {dataset_name} ({target_split})" ) else: logger.warning(f"Split '{target_split}' not found in {dataset_name}. Available: {list(ds.keys())}") except Exception as e: logger.error(f"Error loading {dataset_name}: {str(e)}") continue logger.info(f"Total classification examples: {len(all_data)}") return all_data def load_ner_dataset( self, datasets: List[str] = ["cord", "funsd"], split: str = "train" ) -> List[Dict]: """ Load and combine NER data from multiple datasets """ all_data = [] for dataset_name in datasets: if dataset_name not in self.loaders: logger.warning(f"Unknown dataset: {dataset_name}") continue try: loader = self.loaders[dataset_name] ds = loader.load_dataset_splits() # Handle different split names target_split = split if split == "validation": if "val" in ds: target_split = "val" elif "dev" in ds: target_split = "dev" elif "test" in ds: target_split = "test" if target_split in ds: data = loader.extract_ner_data(ds[target_split]) all_data.extend(data) logger.info( f"Loaded {len(data)} NER examples from {dataset_name} ({target_split})" ) else: logger.warning( f"Split '{target_split}' not found in {dataset_name}" ) except Exception as e: logger.error(f"Error loading {dataset_name}: {str(e)}") continue logger.info(f"Total NER examples: {len(all_data)}") return all_data def get_label_mappings(self): """Get label mappings for classification and NER""" # Classification labels classification_labels = ["INVOICE", "RECEIPT", "FORM", "OTHER"] # NER labels (BIO tagging) ner_labels = [ "O", "B-INVOICE_NUMBER", "I-INVOICE_NUMBER", "B-DATE", "I-DATE", "B-TOTAL_AMOUNT", "I-TOTAL_AMOUNT", "B-TAX_AMOUNT", "I-TAX_AMOUNT", "B-VENDOR_NAME", "I-VENDOR_NAME", "B-CUSTOMER_NAME", "I-CUSTOMER_NAME", "B-ADDRESS", "I-ADDRESS", "B-GST_ID", "I-GST_ID", "B-HEADER", "I-HEADER", # FUNSD "B-QUESTION", "I-QUESTION", "B-ANSWER", "I-ANSWER" ] return { "classification": { label: idx for idx, label in enumerate(classification_labels) }, "ner": {label: idx for idx, label in enumerate(ner_labels)}, "classification_id2label": { idx: label for idx, label in enumerate(classification_labels) }, "ner_id2label": {idx: label for idx, label in enumerate(ner_labels)}, } if __name__ == "__main__": # Example usage loader = UnifiedDatasetLoader() # Load classification data print("\n" + "=" * 50) print("Loading Classification Data") print("=" * 50) classification_data = loader.load_classification_dataset( datasets=["cord", "sroie"], split="train" ) print(f"\nTotal examples: {len(classification_data)}") if classification_data: print(f"\nExample:") print(f"Label: {classification_data[0]['label']}") print(f"Text preview: {classification_data[0]['text'][:200]}...") # Load NER data print("\n" + "=" * 50) print("Loading NER Data") print("=" * 50) ner_data = loader.load_ner_dataset(datasets=["cord", "sroie"], split="train") print(f"\nTotal examples: {len(ner_data)}") if ner_data: print(f"\nExample:") print(f"Tokens: {ner_data[0]['tokens'][:10]}...") print(f"Labels: {ner_data[0]['labels'][:10]}...") # Get label mappings print("\n" + "=" * 50) print("Label Mappings") print("=" * 50) mappings = loader.get_label_mappings() print(f"\nClassification labels: {list(mappings['classification'].keys())}") print(f"NER labels: {list(mappings['ner'].keys())[:10]}...")