| """
|
| 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_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,
|
| (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",
|
| num_heads=num_heads,
|
| hidden_size=hidden_size,
|
| opt_level=99,
|
| use_gpu=True,
|
| only_onnxruntime=False,
|
| )
|
|
|
| print("[export] converting to FP16 ...")
|
| optimized.convert_float_to_float16(
|
| keep_io_types=True,
|
| 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)
|
|
|
|
|
| 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}")
|
|
|
|
|
| onnx_path = output_dir / "cross_encoder_base.onnx"
|
| export_to_onnx(model, onnx_path)
|
|
|
|
|
| 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)
|
|
|