如何从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
相关产品推荐
相关产品推荐

