IDP-Machine-learning / ner_model.py
mrrobot2610's picture
Initial commit: IDP (Intelligent Document Processing) System
1a7ee60
Raw History Blame Contribute Delete
12.9 kB
"""
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}")