""" Lightweight Document Classifier using MiniLM Fine-tuned for invoice, receipt, and form classification Optimized for CPU inference """ import torch import torch.nn as nn from transformers import AutoTokenizer, AutoModel, AutoConfig from typing import Dict, List, Tuple import logging import numpy as np logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) class DocumentClassifier(nn.Module): """ Lightweight document classifier based on MiniLM 4 classes: INVOICE, RECEIPT, FORM, OTHER """ def __init__( self, model_name: str = "nreimers/MiniLM-L6-H384-uncased", num_labels: int = 4, dropout_prob: float = 0.1 ): super().__init__() self.num_labels = num_labels self.model_name = model_name # Load pre-trained MiniLM self.config = AutoConfig.from_pretrained(model_name) self.backbone = AutoModel.from_pretrained(model_name, config=self.config) # Classification head self.dropout = nn.Dropout(dropout_prob) self.classifier = nn.Linear(self.config.hidden_size, num_labels) logger.info(f"Initialized DocumentClassifier with {model_name}") logger.info(f"Model parameters: {sum(p.numel() for p in self.parameters()) / 1e6:.2f}M") def forward( self, input_ids: torch.Tensor, attention_mask: torch.Tensor, labels: torch.Tensor = None, **kwargs ) -> Dict[str, torch.Tensor]: """ Forward pass Args: input_ids: Token IDs (batch_size, seq_len) attention_mask: Attention mask (batch_size, seq_len) labels: Ground truth labels (batch_size,) Returns: Dictionary with loss (if labels provided) and logits """ # Get embeddings from backbone outputs = self.backbone( input_ids=input_ids, attention_mask=attention_mask ) # Use [CLS] token representation pooled_output = outputs.last_hidden_state[:, 0, :] # (batch_size, hidden_size) # Apply dropout and classification pooled_output = self.dropout(pooled_output) logits = self.classifier(pooled_output) # (batch_size, num_labels) # Calculate loss if labels provided loss = None if labels is not None: loss_fct = nn.CrossEntropyLoss() loss = loss_fct(logits, labels) return { "loss": loss, "logits": logits, } class DocumentClassifierInference: """ Inference wrapper for document classification Handles tokenization and prediction """ def __init__( self, model_path: str, model_name: str = "nreimers/MiniLM-L6-H384-uncased", device: str = None, use_fp16: bool = False ): """ Args: model_path: Path to fine-tuned model weights model_name: Base model name (for tokenizer) device: Device to run on ('cpu', 'cuda', or None for auto) use_fp16: Use FP16 precision (faster on GPU) """ self.device = device if device else ('cuda' if torch.cuda.is_available() else 'cpu') logger.info(f"Loading model on device: {self.device}") # Load tokenizer self.tokenizer = AutoTokenizer.from_pretrained(model_name) # Load model self.model = DocumentClassifier(model_name=model_name) # Load fine-tuned weights try: state_dict = torch.load(model_path, map_location=self.device) self.model.load_state_dict(state_dict) logger.info(f"Loaded fine-tuned weights from {model_path}") except Exception as e: logger.warning(f"Could not load weights from {model_path}: {str(e)}") logger.warning("Using pre-trained weights (not fine-tuned)") self.model.to(self.device) self.model.eval() # Apply FP16 if requested if use_fp16 and self.device == 'cuda': self.model.half() logger.info("Using FP16 precision") # Label mapping self.id2label = { 0: 'INVOICE', 1: 'RECEIPT', 2: 'FORM', 3: 'OTHER' } self.label2id = {v: k for k, v in self.id2label.items()} @torch.no_grad() def predict( self, text: str, max_length: int = 512, return_probabilities: bool = True ) -> Dict: """ Predict document class from text Args: text: Input text (OCR extracted) max_length: Maximum sequence length return_probabilities: Return class probabilities Returns: Dictionary with predicted class and probabilities """ # Tokenize inputs = self.tokenizer( text, max_length=max_length, padding='max_length', truncation=True, return_tensors='pt' ) # Move to device inputs = {k: v.to(self.device) for k, v in inputs.items()} # Forward pass outputs = self.model(**inputs) logits = outputs['logits'] # Get predictions probabilities = torch.softmax(logits, dim=-1) predicted_class_id = torch.argmax(probabilities, dim=-1).item() predicted_class = self.id2label[predicted_class_id] confidence = probabilities[0, predicted_class_id].item() result = { 'predicted_class': predicted_class, 'confidence': confidence, } if return_probabilities: result['probabilities'] = { self.id2label[i]: probabilities[0, i].item() for i in range(len(self.id2label)) } return result @torch.no_grad() def predict_batch( self, texts: List[str], max_length: int = 512, batch_size: int = 8 ) -> List[Dict]: """ Predict document classes for multiple texts Args: texts: List of input texts max_length: Maximum sequence length batch_size: Batch size for processing Returns: List of prediction dictionaries """ all_results = [] for i in range(0, len(texts), batch_size): batch_texts = texts[i:i+batch_size] # Tokenize batch inputs = self.tokenizer( batch_texts, max_length=max_length, padding='max_length', truncation=True, return_tensors='pt' ) # Move to device inputs = {k: v.to(self.device) for k, v in inputs.items()} # Forward pass outputs = self.model(**inputs) logits = outputs['logits'] # Get predictions probabilities = torch.softmax(logits, dim=-1) predicted_classes = torch.argmax(probabilities, dim=-1) # Parse results for j in range(len(batch_texts)): pred_id = predicted_classes[j].item() result = { 'predicted_class': self.id2label[pred_id], 'confidence': probabilities[j, pred_id].item(), 'probabilities': { self.id2label[k]: probabilities[j, k].item() for k in range(len(self.id2label)) } } all_results.append(result) return all_results def export_to_onnx( model_path: str, output_path: str, model_name: str = "nreimers/MiniLM-L6-H384-uncased" ): """ Export model to ONNX format for faster inference Args: model_path: Path to PyTorch model weights output_path: Path to save ONNX model model_name: Base model name """ import torch.onnx logger.info("Exporting model to ONNX...") # Load model model = DocumentClassifier(model_name=model_name) state_dict = torch.load(model_path, map_location='cpu') model.load_state_dict(state_dict) model.eval() # Load tokenizer tokenizer = AutoTokenizer.from_pretrained(model_name) # Create dummy input dummy_text = "This is a sample invoice text for ONNX export" dummy_input = tokenizer( dummy_text, max_length=512, padding='max_length', truncation=True, return_tensors='pt' ) # Export to ONNX torch.onnx.export( model, (dummy_input['input_ids'], dummy_input['attention_mask']), output_path, input_names=['input_ids', 'attention_mask'], output_names=['logits'], dynamic_axes={ 'input_ids': {0: 'batch_size'}, 'attention_mask': {0: 'batch_size'}, 'logits': {0: 'batch_size'} }, opset_version=14, ) logger.info(f"Model exported to {output_path}") if __name__ == "__main__": # Example usage # Create a sample model (not trained) model = DocumentClassifier() print(f"\nModel architecture:") print(model) print(f"\nTotal parameters: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M") # Test forward pass batch_size = 2 seq_len = 128 dummy_input_ids = torch.randint(0, 1000, (batch_size, seq_len)) dummy_attention_mask = torch.ones(batch_size, seq_len) dummy_labels = torch.tensor([0, 1]) outputs = model(dummy_input_ids, dummy_attention_mask, dummy_labels) print(f"\nOutput logits shape: {outputs['logits'].shape}") print(f"Loss: {outputs['loss'].item():.4f}")