""" Training Script for Document Classifier Fine-tunes MiniLM on CORD, SROIE, and FUNSD datasets Includes data augmentation and early stopping """ import json import logging import os from typing import Dict, List import numpy as np import torch import torch.nn as nn from sklearn.metrics import accuracy_score, classification_report, f1_score from sklearn.model_selection import train_test_split from torch.optim import AdamW from torch.utils.data import DataLoader, Dataset from tqdm import tqdm from transformers import AutoTokenizer, get_linear_schedule_with_warmup from classifier_model import DocumentClassifier from dataset_loader import UnifiedDatasetLoader logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) class DocumentClassificationDataset(Dataset): """PyTorch Dataset for document classification""" def __init__( self, texts: List[str], labels: List[int], tokenizer, max_length: int = 512 ): self.texts = texts self.labels = labels self.tokenizer = tokenizer self.max_length = max_length def __len__(self): return len(self.texts) def __getitem__(self, idx): text = self.texts[idx] label = self.labels[idx] # Tokenize encoding = self.tokenizer( text, max_length=self.max_length, padding="max_length", truncation=True, return_tensors="pt", ) return { "input_ids": encoding["input_ids"].flatten(), "attention_mask": encoding["attention_mask"].flatten(), "labels": torch.tensor(label, dtype=torch.long), } class TextAugmenter: """Simple text augmentation for document classification""" @staticmethod def random_token_masking(text: str, mask_prob: float = 0.1) -> str: """Randomly mask tokens with [MASK]""" tokens = text.split() num_to_mask = int(len(tokens) * mask_prob) if num_to_mask > 0: mask_indices = np.random.choice(len(tokens), num_to_mask, replace=False) for idx in mask_indices: tokens[idx] = "[MASK]" return " ".join(tokens) @staticmethod def random_deletion(text: str, delete_prob: float = 0.1) -> str: """Randomly delete tokens""" tokens = text.split() tokens = [t for t in tokens if np.random.random() > delete_prob] return " ".join(tokens) if tokens else text class ClassifierTrainer: """Trainer for document classifier""" def __init__( self, model_name: str = "nreimers/MiniLM-L6-H384-uncased", output_dir: str = "models/classifier", device: str = None, ): self.model_name = model_name self.output_dir = output_dir self.device = ( device if device else ("cuda" if torch.cuda.is_available() else "cpu") ) os.makedirs(output_dir, exist_ok=True) logger.info(f"Trainer initialized on device: {self.device}") def prepare_data( self, augment: bool = True, test_size: float = 0.1, val_size: float = 0.1 ): """Load and prepare datasets""" logger.info("Loading datasets...") # Load data from all sources loader = UnifiedDatasetLoader() train_data = loader.load_classification_dataset( datasets=["cord", "sroie", "funsd"], split="train" ) # Get label mappings mappings = loader.get_label_mappings() self.label2id = mappings["classification"] self.id2label = mappings["classification_id2label"] # Extract texts and labels texts = [item["text"] for item in train_data] labels = [self.label2id[item["label"]] for item in train_data] # Apply augmentation if augment: logger.info("Applying data augmentation...") augmenter = TextAugmenter() aug_texts = [] aug_labels = [] for text, label in zip(texts, labels): # Original aug_texts.append(text) aug_labels.append(label) # Augmented versions (30% of data) if np.random.random() < 0.3: aug_texts.append(augmenter.random_token_masking(text)) aug_labels.append(label) if np.random.random() < 0.3: aug_texts.append(augmenter.random_deletion(text)) aug_labels.append(label) texts = aug_texts labels = aug_labels logger.info(f"Augmented dataset size: {len(texts)}") # Split into train/val/test train_texts, temp_texts, train_labels, temp_labels = train_test_split( texts, labels, test_size=(test_size + val_size), random_state=42, stratify=labels, ) val_texts, test_texts, val_labels, test_labels = train_test_split( temp_texts, temp_labels, test_size=test_size / (test_size + val_size), random_state=42, stratify=temp_labels, ) logger.info( f"Train: {len(train_texts)}, Val: {len(val_texts)}, Test: {len(test_texts)}" ) return ( (train_texts, train_labels), (val_texts, val_labels), (test_texts, test_labels), ) def train( self, train_data: tuple, val_data: tuple, num_epochs: int = 10, batch_size: int = 16, learning_rate: float = 2e-5, warmup_steps: int = 500, early_stopping_patience: int = 3, ): """Train the classifier""" train_texts, train_labels = train_data val_texts, val_labels = val_data # Load tokenizer and model tokenizer = AutoTokenizer.from_pretrained(self.model_name) model = DocumentClassifier( model_name=self.model_name, num_labels=len(self.label2id) ) model.to(self.device) # Create datasets train_dataset = DocumentClassificationDataset( train_texts, train_labels, tokenizer ) val_dataset = DocumentClassificationDataset(val_texts, val_labels, tokenizer) # Create dataloaders train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=batch_size) # Optimizer and scheduler optimizer = AdamW(model.parameters(), lr=learning_rate) total_steps = len(train_loader) * num_epochs scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_steps ) # Training loop best_val_accuracy = 0 patience_counter = 0 for epoch in range(num_epochs): logger.info(f"\nEpoch {epoch + 1}/{num_epochs}") # Training model.train() train_loss = 0 train_preds = [] train_true = [] progress_bar = tqdm(train_loader, desc="Training") for batch in progress_bar: optimizer.zero_grad() input_ids = batch["input_ids"].to(self.device) attention_mask = batch["attention_mask"].to(self.device) labels = batch["labels"].to(self.device) outputs = model(input_ids, attention_mask, labels) loss = outputs["loss"] loss.backward() optimizer.step() scheduler.step() train_loss += loss.item() preds = torch.argmax(outputs["logits"], dim=-1) train_preds.extend(preds.cpu().numpy()) train_true.extend(labels.cpu().numpy()) progress_bar.set_postfix({"loss": loss.item()}) avg_train_loss = train_loss / len(train_loader) train_accuracy = accuracy_score(train_true, train_preds) train_f1 = f1_score(train_true, train_preds, average="weighted") logger.info( f"Train Loss: {avg_train_loss:.4f}, Accuracy: {train_accuracy:.4f}, F1: {train_f1:.4f}" ) # Validation model.eval() val_loss = 0 val_preds = [] val_true = [] with torch.no_grad(): for batch in tqdm(val_loader, desc="Validation"): input_ids = batch["input_ids"].to(self.device) attention_mask = batch["attention_mask"].to(self.device) labels = batch["labels"].to(self.device) outputs = model(input_ids, attention_mask, labels) val_loss += outputs["loss"].item() preds = torch.argmax(outputs["logits"], dim=-1) val_preds.extend(preds.cpu().numpy()) val_true.extend(labels.cpu().numpy()) avg_val_loss = val_loss / len(val_loader) val_accuracy = accuracy_score(val_true, val_preds) val_f1 = f1_score(val_true, val_preds, average="weighted") logger.info( f"Val Loss: {avg_val_loss:.4f}, Accuracy: {val_accuracy:.4f}, F1: {val_f1:.4f}" ) # Early stopping if val_accuracy > best_val_accuracy: best_val_accuracy = val_accuracy patience_counter = 0 # Save best model model_path = os.path.join(self.output_dir, "best_classifier.pt") torch.save(model.state_dict(), model_path) logger.info(f"Saved best model with accuracy: {best_val_accuracy:.4f}") else: patience_counter += 1 if patience_counter >= early_stopping_patience: logger.info(f"Early stopping triggered after {epoch + 1} epochs") break # Save final model and metadata torch.save( model.state_dict(), os.path.join(self.output_dir, "final_classifier.pt") ) metadata = { "model_name": self.model_name, "num_labels": len(self.label2id), "label2id": self.label2id, "id2label": self.id2label, "best_val_accuracy": best_val_accuracy, } with open(os.path.join(self.output_dir, "metadata.json"), "w") as f: json.dump(metadata, f, indent=2) logger.info("Training complete!") return best_val_accuracy if __name__ == "__main__": # Training configuration trainer = ClassifierTrainer( model_name="nreimers/MiniLM-L6-H384-uncased", output_dir="models/classifier" ) # Prepare data train_data, val_data, test_data = trainer.prepare_data( augment=True, test_size=0.1, val_size=0.1 ) # Train model best_accuracy = trainer.train( train_data=train_data, val_data=val_data, num_epochs=15, batch_size=16, learning_rate=2e-5, early_stopping_patience=3, ) print(f"\nBest validation accuracy: {best_accuracy:.4f}")