ManoVignesh/Invoice_Information_Extraction_using_a_LayoutLMv3
0
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}")