| 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() |
|
|