#!/usr/bin/env python3 """Convert TensorFlow SavedModel to ONNX format.""" import argparse import logging import subprocess import sys from pathlib import Path def main(): """Convert a TensorFlow SavedModel to ONNX format using tf2onnx.""" logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") parser = argparse.ArgumentParser( description="Convert TensorFlow SavedModel to ONNX format" ) parser.add_argument( "--saved-model", type=Path, default=Path("model/saved_model"), help="Path to the TensorFlow SavedModel directory", ) parser.add_argument( "--output", type=Path, default=Path("model/model.onnx"), help="Path to the output ONNX model file", ) parser.add_argument( "--opset", type=int, default=18, help="ONNX opset version to target", ) args = parser.parse_args() if not args.saved_model.exists(): msg = f"SavedModel not found at {args.saved_model}" logging.error(msg) raise FileNotFoundError(msg) args.output.parent.mkdir(parents=True, exist_ok=True) cmd = [ sys.executable, "-m", "tf2onnx.convert", "--saved-model", str(args.saved_model), "--output", str(args.output), "--opset", str(args.opset), ] try: subprocess.run(cmd, check=True, capture_output=True, text=True) except subprocess.CalledProcessError as e: logging.error(f"tf2onnx conversion failed with exit code {e.returncode}") if e.stderr: logging.error(f"stderr: {e.stderr}") raise logging.info(f"ONNX model saved to {args.output}") if __name__ == "__main__": main()