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

如何从tf.compat.v1.get_default_graph()加载的pb计算图获取各层输出形状

针对TensorFlow pb模型获取层输出形状的方法

是可以获取的,具体分为静态形状获取(无需运行模型,拿计算图中预存的推断形状)和动态运行时形状获取(喂入输入数据后拿实际运行的维度)两种方案:

1 获取静态输出形状

pb文件存储的GraphDef结构中,每个运算节点的属性已经记录了编译阶段推断出的静态输出形状,直接读取即可:

  • 第一步先加载pb文件到计算图中,参考代码:
import tensorflow as tf

def load_pb_model(pb_path):
    with tf.io.gfile.GFile(pb_path, "rb") as f:
        graph_def = tf.compat.v1.GraphDef()
        graph_def.ParseFromString(f.read())
    with tf.Graph().as_default() as graph:
        tf.import_graph_def(graph_def, name="")
    return graph
  • 第二步遍历所有运算节点的输出张量,读取shape属性:
# 替换为你的pb文件路径
graph = load_pb_model("./model.pb")
for op in graph.get_operations():
    # 单个运算可能有多个输出张量,遍历全部输出
    for idx, output in enumerate(op.outputs):
        print(f"运算名:{op.name},输出序号:{idx},静态形状:{output.shape}")

如果输出的形状里存在?或者<unknown>,说明该维度是动态的,静态阶段无法确定,需要运行模型喂入输入才能拿到实际值。

2 获取动态运行时形状

对于动态维度的输出,需要构造符合要求的输入数据,运行计算图后获取实际形状:

with tf.compat.v1.Session(graph=graph) as sess:
    # 替换为你模型的输入节点名称,末尾的:0代表第一个输出
    input_tensor = graph.get_tensor_by_name("你的输入节点名:0")
    # 构造符合输入要求的测试数据,示例为batch=1的224*224三通道图片
    test_input = tf.random.normal(shape=[1, 224, 224, 3]).eval(session=sess)
    
    for op in graph.get_operations():
        for idx, output in enumerate(op.outputs):
            # 跳过资源类、控制流类无实际张量输出的节点
            if output.dtype == tf.dtypes.resource:
                continue
            try:
                real_shape = sess.run(tf.shape(output), feed_dict={input_tensor: test_input})
                print(f"运算名:{op.name},输出序号:{idx},运行时形状:{real_shape}")
            except:
                # 存在多输入、控制流依赖的节点无法单独运行的话直接跳过即可
                continue

注意事项

  • 如果模型有多个输入,需要在feed_dict中传入所有要求的输入数据
  • 变量类、初始化类、控制流类节点不属于前向传播的计算层,无有效输出形状可以直接忽略

内容的提问来源于stack exchange,提问作者Dittsche

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 07:30:00