You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.29 08:05:17