""" Lightweight NER Model using DistilBERT Token-level entity extraction for document fields Optimized for CPU inference with quantization support """ import torch import torch.nn as nn from transformers import ( AutoTokenizer, AutoModel, AutoConfig, DistilBertModel, DistilBertConfig ) from typing import Dict, List, Tuple, Optional import logging import numpy as np logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) class DocumentNERModel(nn.Module): """ Token classification model for NER Based on DistilBERT for lightweight inference """ def __init__( self, model_name: str = "distilbert-base-uncased", num_labels: int = 17, # Number of NER tags (including O and BIO tags) dropout_prob: float = 0.3 ): super().__init__() self.num_labels = num_labels self.model_name = model_name # Load pre-trained DistilBERT self.config = DistilBertConfig.from_pretrained(model_name) self.backbone = DistilBertModel.from_pretrained(model_name, config=self.config) # Token classification head self.dropout = nn.Dropout(dropout_prob) self.classifier = nn.Linear(self.config.hidden_size, num_labels) logger.info(f"Initialized DocumentNERModel 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 ) -> 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, seq_len) Returns: Dictionary with loss (if labels provided) and logits """ # Get token representations outputs = self.backbone( input_ids=input_ids, attention_mask=attention_mask ) # Get last hidden state sequence_output = outputs.last_hidden_state # (batch_size, seq_len, hidden_size) # Apply dropout and classification sequence_output = self.dropout(sequence_output) logits = self.classifier(sequence_output) # (batch_size, seq_len, num_labels) # Calculate loss if labels provided loss = None if labels is not None: loss_fct = nn.CrossEntropyLoss(ignore_index=-100) # Ignore padding tokens loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1)) return { "loss": loss, "logits": logits, } class DocumentNERInference: """ Inference wrapper for NER model Handles tokenization, prediction, and entity extraction """ def __init__( self, model_path: str, model_name: str = "distilbert-base-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 NER model on device: {self.device}") # Load tokenizer self.tokenizer = AutoTokenizer.from_pretrained(model_name) # NER label mappings (BIO tagging) self.id2label = { 0: 'O', 1: 'B-INVOICE_NUMBER', 2: 'I-INVOICE_NUMBER', 3: 'B-DATE', 4: 'I-DATE', 5: 'B-TOTAL_AMOUNT', 6: 'I-TOTAL_AMOUNT', 7: 'B-TAX_AMOUNT', 8: 'I-TAX_AMOUNT', 9: 'B-VENDOR_NAME', 10: 'I-VENDOR_NAME', 11: 'B-CUSTOMER_NAME', 12: 'I-CUSTOMER_NAME', 13: 'B-ADDRESS', 14: 'I-ADDRESS', 15: 'B-GST_ID', 16: 'I-GST_ID', } self.label2id = {v: k for k, v in self.id2label.items()} # Load model self.model = DocumentNERModel( model_name=model_name, num_labels=len(self.id2label) ) # 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") @torch.no_grad() def predict( self, text: str, max_length: int = 512, return_token_scores: bool = False ) -> Dict: """ Predict entities from text Args: text: Input text (OCR extracted) max_length: Maximum sequence length return_token_scores: Return per-token scores Returns: Dictionary with extracted entities """ # Tokenize tokenized = self.tokenizer( text, max_length=max_length, padding='max_length', truncation=True, return_tensors='pt', return_offsets_mapping=True ) offset_mapping = tokenized.pop('offset_mapping')[0] # Move to device inputs = {k: v.to(self.device) for k, v in tokenized.items()} # Forward pass outputs = self.model(**inputs) logits = outputs['logits'] # Get predictions predictions = torch.argmax(logits, dim=-1)[0] # (seq_len,) probabilities = torch.softmax(logits, dim=-1)[0] # (seq_len, num_labels) # Extract entities entities = self._extract_entities( text, predictions.cpu().numpy(), probabilities.cpu().numpy(), offset_mapping.cpu().numpy(), tokenized['input_ids'][0].cpu().numpy() ) result = { 'entities': entities, 'text': text } if return_token_scores: token_scores = [] tokens = self.tokenizer.convert_ids_to_tokens(tokenized['input_ids'][0]) for i, (token, pred_id) in enumerate(zip(tokens, predictions)): if token not in ['[PAD]', '[CLS]', '[SEP]']: token_scores.append({ 'token': token, 'label': self.id2label[pred_id.item()], 'confidence': probabilities[i, pred_id].item() }) result['token_scores'] = token_scores return result def _extract_entities( self, text: str, predictions: np.ndarray, probabilities: np.ndarray, offset_mapping: np.ndarray, input_ids: np.ndarray ) -> List[Dict]: """ Extract entity spans from BIO predictions Returns: List of entity dictionaries with text, label, confidence, and position """ entities = [] current_entity = None for idx, pred_id in enumerate(predictions): if input_ids[idx] in [self.tokenizer.pad_token_id, self.tokenizer.cls_token_id, self.tokenizer.sep_token_id]: continue label = self.id2label[pred_id] confidence = probabilities[idx, pred_id] if label == 'O': # Save current entity if exists if current_entity is not None: entities.append(current_entity) current_entity = None elif label.startswith('B-'): # Start new entity if current_entity is not None: entities.append(current_entity) entity_type = label[2:] # Remove 'B-' prefix start, end = offset_mapping[idx] current_entity = { 'entity': entity_type, 'text': text[start:end], 'start': int(start), 'end': int(end), 'confidence': float(confidence), 'token_count': 1 } elif label.startswith('I-'): # Continue current entity if current_entity is not None: entity_type = label[2:] if current_entity['entity'] == entity_type: # Extend entity start, end = offset_mapping[idx] current_entity['end'] = int(end) current_entity['text'] = text[current_entity['start']:current_entity['end']] # Update confidence (average) current_entity['confidence'] = ( current_entity['confidence'] * current_entity['token_count'] + confidence ) / (current_entity['token_count'] + 1) current_entity['token_count'] += 1 # Don't forget last entity if current_entity is not None: entities.append(current_entity) # Clean up for entity in entities: entity.pop('token_count', None) entity['text'] = entity['text'].strip() return entities def export_to_onnx( model_path: str, output_path: str, model_name: str = "distilbert-base-uncased", num_labels: int = 17 ): """ Export NER 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 num_labels: Number of NER labels """ import torch.onnx logger.info("Exporting NER model to ONNX...") # Load model model = DocumentNERModel(model_name=model_name, num_labels=num_labels) 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 with number INV-12345 dated 2025-11-28" 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', 1: 'sequence_length'}, 'attention_mask': {0: 'batch_size', 1: 'sequence_length'}, 'logits': {0: 'batch_size', 1: 'sequence_length'} }, opset_version=14, ) logger.info(f"NER model exported to {output_path}") if __name__ == "__main__": # Example usage # Create a sample model (not trained) model = DocumentNERModel() 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.randint(0, 17, (batch_size, seq_len)) 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}")