Keras模型转frozen_graph.pb后优化遇TypeError: graph_def需为GraphDef proto
Hey there, let's fix this error right away! The core issue here is that you're passing a file path string to input_graph_def in trt.create_inference_graph(), but this parameter expects a parsed GraphDef proto object—not just a path to your .pb file. That's exactly why you're seeing the TypeError: graph_def must be a GraphDef proto message.
Root Cause Breakdown
The create_inference_graph function doesn't handle loading the .pb file for you automatically. You need to first read the binary content of your frozen graph file and parse it into a valid GraphDef structure that TensorFlow can recognize and process.
Fixed Code Implementation
Here's the corrected code that properly loads, parses, and optimizes your frozen graph with TensorRT:
import tensorflow as tf from tensorflow.python.framework import graph_io import tensorflow.contrib.tensorrt as trt # Define paths and output nodes frozen_graph_path = './model/frozen_model.pb' output_names = ['conv2d_59','conv2d_67','conv2d_75'] # Load and parse the frozen graph into a GraphDef object with tf.gfile.GFile(frozen_graph_path, 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) # Run TensorRT optimization with the valid GraphDef trt_graph = trt.create_inference_graph( input_graph_def=graph_def, # Now passing the parsed GraphDef instead of a path outputs=output_names, max_batch_size=1, max_workspace_size_bytes=1 << 25, precision_mode='FP16', minimum_segment_size=50 ) # Save the optimized TensorRT graph graph_io.write_graph(trt_graph, "./model/", "trt_graph.pb", as_text=False)
Additional Tips for Jetson Platform
- Import Issues: You mentioned problems importing
tensorflow.contrib.tensorrtandgraph_io. Ensure you're using the official TensorFlow build for Jetson (pre-installed or from NVIDIA's Jetson ecosystem) — these builds come with TensorRT integration pre-configured, which avoids compatibility issues with generic pip-installed TensorFlow versions. - Validate Frozen Graph: Before optimization, double-check that your
frozen_model.pbwas correctly converted from the Kerasmodel.h5. You can verify it by loading the graph in TensorBoard or usingtf.get_default_graph().get_operations()to confirm your output nodes (conv2d_59, etc.) exist in the graph.
Original Error Traceback
Traceback (most recent call last):
File "to_tensorrt.py", line 12, in
minimum_segment_size=50
File "/home/christie/yolo_keras/yolo-keras/lib/python3.6/site-packages/tensorflow/contrib/tensorrt/python/trt_convert.py", line 51, in create_inference_graph
session_config=session_config)
File "/home/christie/yolo_keras/yolo-keras/lib/python3.6/site-packages/tensorflow/python/compiler/tensorrt/trt_convert.py", line 1146, in create_inference_graph
converted_graph_def = trt_converter.convert()
File "/home/christie/yolo_keras/yolo-keras/lib/python3.6/site-packages/tensorflow/python/compiler/tensorrt/trt_convert.py", line 298, in convert
self._convert_graph_def()
File "/home/christie/yolo_keras/yolo-keras/lib/python3.6/site-packages/tensorflow/python/compiler/tensorrt/trt_convert.py", line 221, in _convert_graph_def
importer.import_graph_def(self._input_graph_def, name="")
File "/home/christie/yolo_keras/yolo-keras/lib/python3.6/site-packages/tensorflow/python/util/deprecation.py", line 507, in new_func
return func(*args, **kwargs)
File "/home/christie/yolo_keras/yolo-keras/lib/python3.6/site-packages/tensorflow/python/framework/importer.py", line 394, in import_graph_def
graph_def = _ProcessGraphDefParam(graph_def, op_dict)
File "/home/christie/yolo_keras/yolo-keras/lib/python3.6/site-packages/tensorflow/python/framework/importer.py", line 96, in _ProcessGraphDefParam
raise TypeError('graph_def must be a GraphDef proto.')
TypeError: graph_def must be a GraphDef proto.
内容的提问来源于stack exchange,提问作者Sai Krishnadas

