Team Ai
Apppublic

ManoVignesh/Invoice_Information_Extraction_using_a_LayoutLMv3

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
model.py446 linesDownload Raw Back to root
1import sqlite32import json3import os4import numpy as np5from sklearn.model_selection import train_test_split6from transformers import LayoutLMv3Processor, LayoutLMv3ForTokenClassification, TrainingArguments, Trainer7from datasets import Dataset, Features, Sequence, Value, Array2D8from PIL import Image9import torch10from seqeval.metrics import classification_report, accuracy_score11import re12import warnings13import pytesseract14from pytesseract import Output15from dateutil import parser16 17warnings.filterwarnings("ignore", category=FutureWarning)18warnings.filterwarnings("ignore", category=UserWarning)19 20class Config:21    IMAGE_DIR = "invoices/images"22    ANNOTATION_DIR = "invoices/Annotations/layoutlm_HF_format"23    DB_PATH = "extracted_invoices.db"24    25    MODEL_NAME = "microsoft/layoutlmv3-base"26    27    LABEL_NAMES = [28        "O",           # 029        "B-TOTAL",     # 130        "I-TOTAL",     # 231        "B-DATE",      # 332        "I-DATE",      # 433        "B-BUYER",     # 534        "I-BUYER",     # 635        "B-TAX",       # 736        "I-TAX",       # 837        "B-INVOICE",   # 938        "I-INVOICE",   # 1039        "B-STRUCTURE", # 1140        "I-STRUCTURE", # 1241        "OTHER"        # 1342    ]43    NUM_LABELS = len(LABEL_NAMES)44    ID2LABEL = {i: label for i, label in enumerate(LABEL_NAMES)}45    LABEL2ID = {label: i for i, label in enumerate(LABEL_NAMES)}46    47    # Enhanced training parameters48    BATCH_SIZE = 449    LEARNING_RATE = 1e-550    NUM_EPOCHS = 10  # Increased epochs51    WEIGHT_DECAY = 0.152    MAX_SEQ_LENGTH = 51253    TEST_SIZE = 0.354    DROPOUT_RATE = 0.455 56class InvoiceDataset:57    def __init__(self, annotation_dir, image_dir):58        self.annotation_dir = annotation_dir59        self.image_dir = image_dir60        self.processor = LayoutLMv3Processor.from_pretrained(Config.MODEL_NAME, apply_ocr=False)61 62    def load_data(self, limit=50):63        samples = []64        annotation_files = [f for f in os.listdir(self.annotation_dir) if f.endswith('.json')][:limit]65        66        for ann_file in annotation_files:67            with open(os.path.join(self.annotation_dir, ann_file)) as f:68                data = json.load(f)69                70            if any(s['image_path'] == data['path'] for s in samples):71                continue72                73            image_path = os.path.join(self.image_dir, data['path'])74            image = Image.open(image_path).convert("RGB")75            76            # Validate label indices77            valid_ner_tags = []78            for tag in data['ner_tags']:79                if tag >= Config.NUM_LABELS or tag < 0:80                    valid_ner_tags.append(Config.LABEL2ID["OTHER"])81                else:82                    valid_ner_tags.append(tag)83 84            encoding = self.processor(85                image,86                data['words'],87                boxes=data['bboxes'],88                word_labels=valid_ner_tags,89                truncation=True,90                padding="max_length",91                max_length=Config.MAX_SEQ_LENGTH,92                return_offsets_mapping=True,93                return_tensors="pt"94            )95            96            samples.append({97                'id': data['path'],98                'input_ids': encoding['input_ids'].squeeze(),99                'attention_mask': encoding['attention_mask'].squeeze(),100                'bbox': encoding['bbox'].squeeze(),101                'labels': encoding['labels'].squeeze(),102                'image_path': image_path103            })104        return samples105 106def split_dataset(dataset):107    # Create stratification labels108    labels = [109        1 if any(110            tag in [111                Config.LABEL2ID["B-TOTAL"], 112                Config.LABEL2ID["B-DATE"],113                Config.LABEL2ID["B-BUYER"], 114                Config.LABEL2ID["B-TAX"]115            ] 116            for tag in sample['labels']117        ) 118        else 0 119        for sample in dataset120    ]121    122    # First split: 80% train+val, 20% test123    train_val_data, test_data = train_test_split(124        dataset,125        test_size=0.2,126        stratify=labels,127        random_state=42128    )129    130    # Get labels for train_val subset131    train_val_indices = [i for i, sample in enumerate(dataset) if sample in train_val_data]132    train_val_labels = [labels[i] for i in train_val_indices]133    134    # Second split: 60% train, 20% val135    train_data, val_data = train_test_split(136        train_val_data,137        test_size=0.25,138        stratify=train_val_labels,139        random_state=42140    )141    142    return train_data, val_data, test_data  # Now returns 3 datasets!143 144class InvoiceModelTrainer:145    def __init__(self):146        self.processor = LayoutLMv3Processor.from_pretrained(Config.MODEL_NAME)147        self.model = LayoutLMv3ForTokenClassification.from_pretrained(148            Config.MODEL_NAME,149            num_labels=Config.NUM_LABELS,150            id2label=Config.ID2LABEL,151            label2id=Config.LABEL2ID,152            hidden_dropout_prob=Config.DROPOUT_RATE,153            attention_probs_dropout_prob=Config.DROPOUT_RATE154        )155 156    def compute_metrics(self, p):157        predictions, labels = p158        predictions = np.argmax(predictions, axis=2)159 160        true_predictions = []161        true_labels = []162        163        for prediction, label in zip(predictions, labels):164            valid_preds = []165            valid_lbls = []166            167            for p, l in zip(prediction, label):168                if l != -100 and Config.ID2LABEL.get(p, "O") != "O":169                    valid_preds.append(Config.ID2LABEL[p])170                    valid_lbls.append(Config.ID2LABEL[l])171            172            true_predictions.append(valid_preds)173            true_labels.append(valid_lbls)174 175        return {176            "accuracy": accuracy_score(true_labels, true_predictions),177            **classification_report(true_labels, true_predictions, output_dict=True)178        }179 180    def train(self, train_data, val_data):  # Now takes val_data instead of test_data181        features = Features({182            'input_ids': Sequence(Value(dtype='int64')),183            'attention_mask': Sequence(Value(dtype='int64')),184            'bbox': Array2D(dtype="int64", shape=(Config.MAX_SEQ_LENGTH, 4)),185            'labels': Sequence(Value(dtype='int64'))186        })187        188        # Convert all datasets189        train_dataset = Dataset.from_dict({190            'input_ids': [s['input_ids'].tolist() for s in train_data],191            'attention_mask': [s['attention_mask'].tolist() for s in train_data],192            'bbox': [s['bbox'].tolist() for s in train_data],193            'labels': [s['labels'].tolist() for s in train_data]194        }, features=features)195 196        val_dataset = Dataset.from_dict({197            'input_ids': [s['input_ids'].tolist() for s in val_data],198            'attention_mask': [s['attention_mask'].tolist() for s in val_data],199            'bbox': [s['bbox'].tolist() for s in val_data],200            'labels': [s['labels'].tolist() for s in val_data]201        }, features=features)202 203        training_args = TrainingArguments(204            output_dir="./results",205            num_train_epochs=Config.NUM_EPOCHS,206            per_device_train_batch_size=Config.BATCH_SIZE,207            per_device_eval_batch_size=Config.BATCH_SIZE,208            learning_rate=Config.LEARNING_RATE,209            weight_decay=Config.WEIGHT_DECAY,210            evaluation_strategy="epoch",211            save_strategy="epoch",212            logging_dir='./logs',213            load_best_model_at_end=True,214            save_total_limit=2,215            metric_for_best_model="eval_loss",216            greater_is_better=False,217        )218 219        trainer = Trainer(220            model=self.model,221            args=training_args,222            train_dataset=train_dataset,223            eval_dataset=val_dataset,224            compute_metrics=self.compute_metrics,225        )226 227        trainer.train()228        return trainer # Return trained model and metrics229 230class InvoiceProcessor:231    def __init__(self, model_path):232        self.processor = LayoutLMv3Processor.from_pretrained(Config.MODEL_NAME, apply_ocr=False)233        self.model = LayoutLMv3ForTokenClassification.from_pretrained(model_path)234        self.conn = sqlite3.connect(Config.DB_PATH)235        self._init_db()236 237    def _init_db(self):238        cursor = self.conn.cursor()239        cursor.execute('''240        CREATE TABLE IF NOT EXISTS invoices (241            id INTEGER PRIMARY KEY AUTOINCREMENT,242            date TEXT,243            buyer TEXT,244            total REAL,245            tax REAL,246            image_path TEXT247        )''')248        self.conn.commit()249 250    def process_invoice(self, image_path):251        image = Image.open(image_path).convert("RGB")252        253        # Improved OCR with PSM 6254        ocr_data = pytesseract.image_to_data(255            image,256            config='--psm 6 --oem 3',257            output_type=Output.DICT258        )259        260        words = []261        boxes = []262        for i in range(len(ocr_data['text'])):263            text = ocr_data['text'][i].strip()264            if text:265                x = ocr_data['left'][i]266                y = ocr_data['top'][i]267                w = ocr_data['width'][i]268                h = ocr_data['height'][i]269                boxes.append([x, y, x + w, y + h])270                words.append(text)271 272        # Validate box normalization273        image_width, image_height = image.size274        normalized_boxes = []275        for box in boxes:276            if image_width == 0 or image_height == 0:277                continue  # Prevent division by zero278            normalized_box = [279                max(0, min(1000, int(1000 * (box[0] / image_width)))),280                max(0, min(1000, int(1000 * (box[1] / image_height)))),281                max(0, min(1000, int(1000 * (box[2] / image_width)))),282                max(0, min(1000, int(1000 * (box[3] / image_height)))),283            ]284            normalized_boxes.append(normalized_box)285 286        encoding = self.processor(287            image,288            words,289            boxes=normalized_boxes,290            return_tensors="pt",291            truncation=True,292            padding="max_length",293            max_length=Config.MAX_SEQ_LENGTH,294            return_offsets_mapping=True,295        )296 297        model_inputs = {298            "input_ids": encoding["input_ids"],299            "attention_mask": encoding["attention_mask"],300            "bbox": encoding["bbox"],301        }302 303        with torch.no_grad():304            outputs = self.model(**model_inputs)305        predictions = outputs.logits.argmax(-1).squeeze().tolist()306 307        offset_mapping = encoding['offset_mapping'].squeeze().tolist()308        word_predictions = []309        current_word_id = -1310 311        for i, (offset, pred) in enumerate(zip(offset_mapping, predictions)):312            if offset == (0, 0):313                continue314            315            if i == 0 or offset[0] != offset_mapping[i-1][1]:316                current_word_id += 1317                if current_word_id >= len(words):318                    continue319                word_predictions.append(Config.ID2LABEL.get(pred, "O"))320 321        # Fixed entity extraction logic322        entities = {"date": "", "buyer": "", "total": "", "tax": ""}323        current_entity = None324        current_value = []325        326        for word, label in zip(words, word_predictions):327            if label.startswith("B-"):328                if current_entity:329                    entities[current_entity] = " ".join(current_value)330                    current_value = []331                current_entity = label.split("-")[1].lower()332                current_value.append(word)333            elif label.startswith("I-") and current_entity:334                current_value.append(word)335            else:336                if current_entity:337                    entities[current_entity] = " ".join(current_value)338                    current_entity = None339                    current_value = []340        341        if current_entity:342            entities[current_entity] = " ".join(current_value)343 344        # Enhanced validation345        entities["total"] = self._validate_numeric(entities["total"])346        entities["tax"] = self._validate_numeric(entities["tax"])347        entities["date"] = self._validate_date(entities["date"])348        349        return entities350 351    def _validate_numeric(self, value):352        try:353            cleaned = re.sub(r'[^\d.]', '', str(value))354            return float(cleaned) if cleaned else 0.0355        except:356            return 0.0357 358    def _validate_date(self, value):359        try:360            return parser.parse(value, fuzzy=True).strftime("%d-%b-%Y")361        except:362            return ""363 364    def save_to_db(self, entities, image_path):365        cursor = self.conn.cursor()366        cursor.execute('''INSERT INTO invoices (date, buyer, total, tax, image_path)367                          VALUES (?, ?, ?, ?, ?)''',368                       (entities.get('date', ''),369                        entities.get('buyer', ''),370                        entities.get('total', 0.0),371                        entities.get('tax', 0.0),372                        image_path))373        self.conn.commit()374 375    def print_db_contents(self):376        cursor = self.conn.cursor()377        cursor.execute("SELECT * FROM invoices")378        rows = cursor.fetchall()379        380        print("\n***** Extracted Invoices Database Contents *****")381        print("ID | Date       | Buyer           | Total   | Tax   | Image Path")382        print("-" * 70)383        for row in rows:384            print(f"{row[0]:<3} | {row[1]:<10} | {row[2]:<15} | {row[3]:<7.2f} | {row[4]:<5.2f} | {row[5]}")385 386# Check if the script is being executed directly (not imported as a module)387if __name__ == "__main__":388    # Load invoice dataset with duplicate check389    dataset = InvoiceDataset(Config.ANNOTATION_DIR, Config.IMAGE_DIR).load_data()390    391    # Verify dataset diversity by counting unique invoice IDs392    unique_templates = len(set([sample['id'] for sample in dataset]))393    print(f"Loaded {len(dataset)} invoices ({unique_templates} unique templates)")394 395    # Split dataset into three parts (train/val/test)396    train_data, val_data, test_data = split_dataset(dataset)397 398    # Verify no data leakage between any sets399    train_ids = {sample['id'] for sample in train_data}400    val_ids = {sample['id'] for sample in val_data}401    test_ids = {sample['id'] for sample in test_data}402    403    assert train_ids.isdisjoint(val_ids), "Data leakage: Train/Val overlap detected!"404    assert train_ids.isdisjoint(test_ids), "Data leakage: Train/Test overlap detected!"405    assert val_ids.isdisjoint(test_ids), "Data leakage: Val/Test overlap detected!"406 407    # Initialize and train the model408    model_trainer = InvoiceModelTrainer()409    huggingface_trainer = model_trainer.train(train_data, val_data)  # Get actual Trainer instance410 411    # Save the best model412    best_model_path = "./results/best_model"413    huggingface_trainer.save_model(best_model_path)414 415    # Prepare test dataset for final evaluation416    test_features = Features({417        'input_ids': Sequence(Value(dtype='int64')),418        'attention_mask': Sequence(Value(dtype='int64')),419        'bbox': Array2D(dtype="int64", shape=(Config.MAX_SEQ_LENGTH, 4)),420        'labels': Sequence(Value(dtype='int64'))421    })422    423    test_dataset = Dataset.from_dict({424        'input_ids': [sample['input_ids'].tolist() for sample in test_data],425        'attention_mask': [sample['attention_mask'].tolist() for sample in test_data],426        'bbox': [sample['bbox'].tolist() for sample in test_data],427        'labels': [sample['labels'].tolist() for sample in test_data]428    }, features=test_features)429 430    # Final evaluation on test set431    test_results = huggingface_trainer.evaluate(test_dataset)432    print("\n***** Final Test Set Evaluation *****")433    print(f"Accuracy: {test_results['eval_accuracy'] * 100:.1f}%")434    print(f"Precision: {test_results['eval_weighted avg']['precision'] * 100:.1f}%")435    print(f"Recall: {test_results['eval_weighted avg']['recall'] * 100:.1f}%")436    print(f"F1-score: {test_results['eval_weighted avg']['f1-score'] * 100:.1f}%")437 438    # Process and store test samples439    processor = InvoiceProcessor(best_model_path)440    for sample in test_data[:50]:441        entities = processor.process_invoice(sample['image_path'])442        processor.save_to_db(entities, sample['image_path'])443 444    # Display database contents445    processor.print_db_contents()446    print(f"\nData saved to {Config.DB_PATH}")