如何在TensorFlow中加载YOLO冻结PB模型,添加操作后保存新冻结模型?
如何将YOLO转换后的TensorFlow模型保存为直接输出检测框的.pb文件
我明白你的需求了——你已经把YOLO转成了可用的.pb模型,但每次使用都得先跑模型输出,再手动调用转换函数才能拿到检测框信息,想把这两步合并成一个新的.pb模型,加载后直接输出boxes、scores、classes对吧?其实就是把后处理操作整合到现有计算图里,再重新冻结成独立的模型文件,下面给你一步步说怎么改:
1. 给输出张量添加明确的命名
你当前得到的boxes、scores、classes是匿名张量,导出模型时需要给它们指定清晰的名字,方便后续加载识别。在你的代码里,拿到这三个张量后,用tf.identity()给它们加上名称:
boxes, scores, classes = get_boxes_from_output(l_output, anchors, len(class_names), input_image_shape, score_threshold=score, iou_threshold=iou) # 给输出张量命名,便于后续导出和加载 boxes = tf.identity(boxes, name='detected_boxes') scores = tf.identity(scores, name='detected_scores') classes = tf.identity(classes, name='detected_classes')
另外建议给你的输入占位符也加个名字,避免后续加载时混淆,比如:
input_image_shape = tf.placeholder(dtype=tf.float32,shape=(2, ), name='input_image_shape') training = tf.placeholder(tf.bool, name='training')
2. 冻结计算图(将变量转为常量)
你的计算图现在包含了导入的预训练YOLO模型和后处理操作,需要把所有变量转换成常量,这样导出的.pb文件才是独立可运行的。在创建会话后,添加这段代码:
# 定义要导出的输出节点名称 output_node_names = ['detected_boxes', 'detected_scores', 'detected_classes'] # 将图中的所有变量转换为常量,生成冻结的图定义 frozen_graph_def = tf.graph_util.convert_variables_to_constants( sess, sess.graph_def, output_node_names=output_node_names )
3. 保存新的冻结模型
最后把冻结后的图定义写入新的.pb文件:
# 设置新模型的保存路径 output_model_path = 'model_data/yolo_final.pb' with tf.gfile.GFile(output_model_path, 'wb') as f: f.write(frozen_graph_def.SerializeToString()) print(f"整合后模型已保存到:{output_model_path}")
验证加载新模型的方式
之后你加载这个新模型时,就可以直接获取这三个输出节点,无需再调用转换函数:
new_model_path = 'model_data/yolo_final.pb' detection_graph = tf.Graph() with detection_graph.as_default(): graph_def = tf.GraphDef() with tf.gfile.GFile(new_model_path, 'rb') as fid: graph_def.ParseFromString(fid.read()) tf.import_graph_def(graph_def, name='') # 获取输入和输出节点 l_input = detection_graph.get_tensor_by_name('input_1:0') input_image_shape = detection_graph.get_tensor_by_name('input_image_shape:0') training = detection_graph.get_tensor_by_name('training:0') out_boxes = detection_graph.get_tensor_by_name('detected_boxes:0') out_scores = detection_graph.get_tensor_by_name('detected_scores:0') out_classes = detection_graph.get_tensor_by_name('detected_classes:0') # 运行会话直接得到检测结果 with tf.Session(graph=detection_graph) as sess: out_boxes_result, out_scores_result, out_classes_result = sess.run( [out_boxes, out_scores, out_classes], feed_dict={ l_input: image_data, input_image_shape: [image.size[1], image.size[0]], training: False } )
注意事项
- 确保
get_boxes_from_output函数里的所有操作都是TensorFlow张量操作,不能包含numpy的非张量运算(你说函数正常运行,应该已经满足这点) - 以上代码基于你当前使用的TensorFlow 1.x版本编写,如果你后续切换到TF2.x,API会有差异,需要调整
内容的提问来源于stack exchange,提问作者f10w
相关产品推荐
相关产品推荐

