import argparse import datetime import onnx import tensorflow as tf import tf2onnx def load_graph_def(pb_path): with tf.io.gfile.GFile(pb_path, "rb") as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) return graph_def def main(): parser = argparse.ArgumentParser(description="Export tensorflow_inception_graph.pb to ONNX") parser.add_argument("--pb", default="../pb/tensorflow_inception_graph.pb") parser.add_argument("--opset", type=int, default=18) args = parser.parse_args() graph_def = load_graph_def(args.pb) model_proto, _ = tf2onnx.convert.from_graph_def( graph_def, input_names=["input:0"], output_names=["softmax2:0"], inputs_as_nchw=["input:0"], opset=args.opset, shape_override={"input:0": [1, 224, 224, 3]}, ) onnx.checker.check_model(model_proto) stamp = datetime.datetime.now().strftime("%Y%b").lower() onnx_path = "tensorflow_inception_graph_%s.onnx" % stamp with open(onnx_path, "wb") as f: f.write(model_proto.SerializeToString()) print("wrote", onnx_path) if __name__ == "__main__": main()