Document Question Answering
Transformers
PyTorch
English
document-processing
ocr
ner
text-classification
information-extraction
invoice
receipt
form
Instructions to use mrrobot2610/IDP-Machine-learning with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use mrrobot2610/IDP-Machine-learning with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("document-question-answering", model="mrrobot2610/IDP-Machine-learning")# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("mrrobot2610/IDP-Machine-learning", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download classifier_model.py from mrrobot2610/IDP-Machine-learning: direct link, hf CLI and curl.
- Browser
- Download file 10.4 kB
-
https://huggingface.co/mrrobot2610/IDP-Machine-learning/resolve/main/classifier_model.py
- Command line
-
hf download hf://mrrobot2610/IDP-Machine-learning/classifier_model.py
-
curl -L -o classifier_model.py https://huggingface.co/mrrobot2610/IDP-Machine-learning/resolve/main/classifier_model.py
10.4 kB
| """ | |
| 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()} | |
| 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 | |
| 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}") | |