""" bert/export.py ONNX export + ORT FP16 graph optimization for the Cross-Encoder. Pipeline: 1. Export PyTorch model → ONNX with dynamic axes (batch + sequence length) 2. ORT graph-level fusion (LayerNorm + GELU + Attention → fused CUDA kernels) 3. FP16 weight conversion (lossless for entailment; doubles throughput on Ampere+) 4. Save optimized FP16 model for TensorRT ingestion """ from __future__ import annotations import argparse from pathlib import Path import torch from transformers import AutoTokenizer from bert.model import BertCrossEncoderVerifier, CrossEncoderConfig, build_model def export_to_onnx( model: BertCrossEncoderVerifier, output_path: Path, max_length: int = 512, opset: int = 17, ) -> None: """Export with dynamic batch and sequence axes.""" model.eval() device = next(model.parameters()).device # Dummy inputs for tracing — short sequence, real shapes traced by dynamic axes dummy_input_ids = torch.randint(0, 30000, (1, 128), dtype=torch.long).to(device) dummy_attention_mask = torch.ones((1, 128), dtype=torch.long).to(device) dummy_token_types = torch.zeros((1, 128), dtype=torch.long).to(device) torch.onnx.export( model.model, # underlying HuggingFace model (handles input naming) (dummy_input_ids, dummy_attention_mask, dummy_token_types), str(output_path), input_names=["input_ids", "attention_mask", "token_type_ids"], output_names=["logits"], dynamic_axes={ "input_ids": {0: "batch_size", 1: "sequence_length"}, "attention_mask": {0: "batch_size", 1: "sequence_length"}, "token_type_ids": {0: "batch_size", 1: "sequence_length"}, "logits": {0: "batch_size"}, }, opset_version=opset, do_constant_folding=True, ) size_mb = output_path.stat().st_size / 1_048_576 print(f"[export] ONNX saved → {output_path} ({size_mb:.1f} MB)") def optimize_and_fp16( input_path: Path, output_path: Path, num_heads: int = 12, hidden_size: int = 768, ) -> None: """ ORT graph optimization: fuse attention + LayerNorm + GELU into CUDA kernels, then convert FP32 weights → FP16 (lossless for classification head tasks). """ try: from onnxruntime.transformers.optimizer import optimize_model except ImportError: raise RuntimeError( "onnxruntime-gpu with transformers optimization required.\n" "Install: pip install onnxruntime-gpu" ) print(f"[export] optimizing {input_path} ...") optimized = optimize_model( str(input_path), model_type="bert", # applies to DeBERTa / RoBERTa architectures num_heads=num_heads, hidden_size=hidden_size, opt_level=99, # maximum graph-level fusions use_gpu=True, only_onnxruntime=False, ) print("[export] converting to FP16 ...") optimized.convert_float_to_float16( keep_io_types=True, # keep input/output in FP32 for compatibility min_positive_val=1e-7, max_finite_val=1e4, ) optimized.save_model_to_file(str(output_path)) size_mb = output_path.stat().st_size / 1_048_576 print(f"[export] FP16 optimized model → {output_path} ({size_mb:.1f} MB)") def full_export_pipeline( checkpoint_path: Path, output_dir: Path, backbone: str = "microsoft/deberta-v3-base", device: str = "cuda", ) -> None: output_dir.mkdir(parents=True, exist_ok=True) # 1. Load model config = CrossEncoderConfig(backbone=backbone) model = build_model(config) state = torch.load(checkpoint_path, map_location=device) if "model_state_dict" in state: state = state["model_state_dict"] model.load_state_dict(state) model.to(device).eval() print(f"[export] loaded checkpoint {checkpoint_path}") # 2. ONNX export onnx_path = output_dir / "cross_encoder_base.onnx" export_to_onnx(model, onnx_path) # 3. ORT optimize + FP16 fp16_path = output_dir / "cross_encoder_opt_fp16.onnx" optimize_and_fp16(onnx_path, fp16_path) print(f"\n[export] pipeline complete.") print(f" base ONNX : {onnx_path}") print(f" FP16 ONNX : {fp16_path}") print(f" → feed fp16_path into TRT session (bert/trt_session.py)") if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--output_dir", type=Path, default=Path("onnx")) parser.add_argument("--backbone", type=str, default="microsoft/deberta-v3-base") parser.add_argument("--device", type=str, default="cuda") args = parser.parse_args() full_export_pipeline(args.checkpoint, args.output_dir, args.backbone, args.device)