""" Model Optimization Utilities Convert PyTorch models to ONNX and apply quantization """ import torch import onnx import onnxruntime as ort from transformers import AutoTokenizer import numpy as np import logging import time from typing import Dict, List from pathlib import Path from classifier_model import DocumentClassifier from ner_model import DocumentNERModel logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) class ModelOptimizer: """Optimize models for faster inference""" @staticmethod def convert_classifier_to_onnx( pytorch_model_path: str, output_path: str, model_name: str = "microsoft/MiniLM-L6-H384-uncased", num_labels: int = 4, opset_version: int = 14 ): """ Convert classifier to ONNX format Args: pytorch_model_path: Path to PyTorch weights output_path: Path to save ONNX model model_name: Base model name num_labels: Number of classification labels opset_version: ONNX opset version """ logger.info("Converting classifier to ONNX...") # Load model model = DocumentClassifier(model_name=model_name, num_labels=num_labels) state_dict = torch.load(pytorch_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 for ONNX conversion" dummy_input = tokenizer( dummy_text, max_length=512, padding='max_length', truncation=True, return_tensors='pt' ) # Export 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=opset_version, do_constant_folding=True, ) logger.info(f"Classifier exported to {output_path}") # Verify ModelOptimizer._verify_onnx_model(output_path) @staticmethod def convert_ner_to_onnx( pytorch_model_path: str, output_path: str, model_name: str = "distilbert-base-uncased", num_labels: int = 17, opset_version: int = 14 ): """ Convert NER model to ONNX format Args: pytorch_model_path: Path to PyTorch weights output_path: Path to save ONNX model model_name: Base model name num_labels: Number of NER labels opset_version: ONNX opset version """ logger.info("Converting NER model to ONNX...") # Load model model = DocumentNERModel(model_name=model_name, num_labels=num_labels) state_dict = torch.load(pytorch_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" dummy_input = tokenizer( dummy_text, max_length=512, padding='max_length', truncation=True, return_tensors='pt' ) # Export 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=opset_version, do_constant_folding=True, ) logger.info(f"NER model exported to {output_path}") # Verify ModelOptimizer._verify_onnx_model(output_path) @staticmethod def _verify_onnx_model(onnx_path: str): """Verify ONNX model is valid""" try: onnx_model = onnx.load(onnx_path) onnx.checker.check_model(onnx_model) logger.info(f"ONNX model verified: {onnx_path}") except Exception as e: logger.error(f"ONNX verification failed: {str(e)}") raise @staticmethod def quantize_onnx_model( input_path: str, output_path: str, quantization_mode: str = "IntegerOps" ): """ Apply dynamic quantization to ONNX model Reduces model size and improves CPU inference speed Args: input_path: Path to ONNX model output_path: Path to save quantized model quantization_mode: "IntegerOps" or "QLinearOps" """ from onnxruntime.quantization import quantize_dynamic, QuantType logger.info(f"Quantizing ONNX model: {input_path}") quantize_dynamic( input_path, output_path, weight_type=QuantType.QInt8 ) logger.info(f"Quantized model saved to {output_path}") # Compare sizes original_size = Path(input_path).stat().st_size / (1024 * 1024) quantized_size = Path(output_path).stat().st_size / (1024 * 1024) logger.info(f"Original size: {original_size:.2f} MB") logger.info(f"Quantized size: {quantized_size:.2f} MB") logger.info(f"Size reduction: {(1 - quantized_size/original_size)*100:.1f}%") class ONNXInferenceSession: """ONNX Runtime inference session wrapper""" def __init__(self, model_path: str, providers: List[str] = None): """ Initialize ONNX Runtime session Args: model_path: Path to ONNX model providers: Execution providers (e.g., ['CPUExecutionProvider']) """ if providers is None: providers = ['CPUExecutionProvider'] self.session = ort.InferenceSession(model_path, providers=providers) self.input_names = [inp.name for inp in self.session.get_inputs()] self.output_names = [out.name for out in self.session.get_outputs()] logger.info(f"ONNX session initialized: {model_path}") logger.info(f"Inputs: {self.input_names}") logger.info(f"Outputs: {self.output_names}") def run(self, inputs: Dict[str, np.ndarray]) -> List[np.ndarray]: """ Run inference Args: inputs: Dictionary of input name -> numpy array Returns: List of output arrays """ # Prepare inputs ort_inputs = {name: inputs[name] for name in self.input_names} # Run outputs = self.session.run(self.output_names, ort_inputs) return outputs def benchmark_models( pytorch_model_path: str, onnx_model_path: str, model_type: str = "classifier", num_runs: int = 100 ): """ Benchmark PyTorch vs ONNX inference speed Args: pytorch_model_path: Path to PyTorch model onnx_model_path: Path to ONNX model model_type: "classifier" or "ner" num_runs: Number of benchmark runs """ logger.info(f"Benchmarking {model_type} models...") # Load tokenizer if model_type == "classifier": model_name = "microsoft/MiniLM-L6-H384-uncased" num_labels = 4 ModelClass = DocumentClassifier else: model_name = "distilbert-base-uncased" num_labels = 17 ModelClass = DocumentNERModel tokenizer = AutoTokenizer.from_pretrained(model_name) # Prepare sample input sample_text = "This is a sample invoice with number INV-12345 dated 28/11/2025" inputs = tokenizer( sample_text, max_length=512, padding='max_length', truncation=True, return_tensors='pt' ) # PyTorch model pytorch_model = ModelClass(model_name=model_name, num_labels=num_labels) pytorch_model.load_state_dict(torch.load(pytorch_model_path, map_location='cpu')) pytorch_model.eval() # ONNX model onnx_session = ONNXInferenceSession(onnx_model_path) # Warmup for _ in range(10): with torch.no_grad(): _ = pytorch_model(inputs['input_ids'], inputs['attention_mask']) onnx_inputs = { 'input_ids': inputs['input_ids'].numpy(), 'attention_mask': inputs['attention_mask'].numpy() } _ = onnx_session.run(onnx_inputs) # Benchmark PyTorch pytorch_times = [] for _ in range(num_runs): start = time.time() with torch.no_grad(): _ = pytorch_model(inputs['input_ids'], inputs['attention_mask']) pytorch_times.append(time.time() - start) # Benchmark ONNX onnx_times = [] for _ in range(num_runs): start = time.time() _ = onnx_session.run(onnx_inputs) onnx_times.append(time.time() - start) # Results logger.info("\nBenchmark Results:") logger.info(f"PyTorch - Mean: {np.mean(pytorch_times)*1000:.2f}ms, " f"Std: {np.std(pytorch_times)*1000:.2f}ms") logger.info(f"ONNX - Mean: {np.mean(onnx_times)*1000:.2f}ms, " f"Std: {np.std(onnx_times)*1000:.2f}ms") logger.info(f"Speedup: {np.mean(pytorch_times)/np.mean(onnx_times):.2f}x") if __name__ == "__main__": import sys if len(sys.argv) < 2: print("Usage:") print(" Convert classifier: python model_optimizer.py convert_classifier ") print(" Convert NER: python model_optimizer.py convert_ner ") print(" Quantize: python model_optimizer.py quantize ") print(" Benchmark: python model_optimizer.py benchmark ") sys.exit(1) command = sys.argv[1] if command == "convert_classifier": pytorch_path = sys.argv[2] onnx_path = sys.argv[3] ModelOptimizer.convert_classifier_to_onnx(pytorch_path, onnx_path) elif command == "convert_ner": pytorch_path = sys.argv[2] onnx_path = sys.argv[3] ModelOptimizer.convert_ner_to_onnx(pytorch_path, onnx_path) elif command == "quantize": input_path = sys.argv[2] output_path = sys.argv[3] ModelOptimizer.quantize_onnx_model(input_path, output_path) elif command == "benchmark": pytorch_path = sys.argv[2] onnx_path = sys.argv[3] model_type = sys.argv[4] benchmark_models(pytorch_path, onnx_path, model_type) else: print(f"Unknown command: {command}")