fix: improve ONNX conversion logging
This commit is contained in:
@@ -2,12 +2,16 @@
|
|||||||
"""Convert TensorFlow SavedModel to ONNX format."""
|
"""Convert TensorFlow SavedModel to ONNX format."""
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import logging
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
"""Convert a TensorFlow SavedModel to ONNX format using tf2onnx."""
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="Convert TensorFlow SavedModel to ONNX format"
|
description="Convert TensorFlow SavedModel to ONNX format"
|
||||||
)
|
)
|
||||||
@@ -32,14 +36,11 @@ def main():
|
|||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
# Check that saved model exists
|
|
||||||
if not args.saved_model.exists():
|
if not args.saved_model.exists():
|
||||||
raise FileNotFoundError(f"SavedModel not found at {args.saved_model}")
|
raise FileNotFoundError(f"SavedModel not found at {args.saved_model}")
|
||||||
|
|
||||||
# Create output parent directory if needed
|
|
||||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
# Run tf2onnx conversion
|
|
||||||
cmd = [
|
cmd = [
|
||||||
sys.executable,
|
sys.executable,
|
||||||
"-m",
|
"-m",
|
||||||
@@ -52,9 +53,13 @@ def main():
|
|||||||
str(args.opset),
|
str(args.opset),
|
||||||
]
|
]
|
||||||
|
|
||||||
subprocess.run(cmd, check=True)
|
try:
|
||||||
|
subprocess.run(cmd, check=True)
|
||||||
|
except subprocess.CalledProcessError as e:
|
||||||
|
logging.error(f"tf2onnx conversion failed with exit code {e.returncode}")
|
||||||
|
raise
|
||||||
|
|
||||||
print(f"ONNX model saved to {args.output}")
|
logging.info(f"ONNX model saved to {args.output}")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user