如何统计从.pb文件加载的TensorFlow图中可训练参数总数?
统计从.pb文件加载的TensorFlow图的参数数量
我太懂你这种困扰了——常规统计可训练参数的方法(比如用model.trainable_variables)对从.pb冻结图加载的模型完全没用,毕竟冻结图里所有变量都已经被转换成常量节点了。别慌,我给你分享一个专门针对.pb图的参数统计方案,亲测有效。
核心思路
冻结后的.pb图里,所有训练过的参数都以Const类型的节点存在,我们只需要遍历图中所有节点,筛选出属于参数的Const节点,再计算每个节点的张量大小并累加即可。
完整代码实现
结合你给出的加载图函数,我把统计逻辑整合进去了:
import tensorflow as tf def load_graph(model_file): graph = tf.Graph() graph_def = tf.GraphDef() with open(model_file, "rb") as f: graph_def.ParseFromString(f.read()) with graph.as_default(): tf.import_graph_def(graph_def, name="") # 补充导入图的关键步骤 return graph def count_pb_model_params(graph): total_params = 0 # 遍历图中所有节点 for node in graph.as_graph_def().node: # 先筛选出常量节点(冻结图的参数都在Const节点里) if node.op == "Const": # 根据节点名关键词判断是否为参数节点,可根据你的模型命名习惯调整 param_keywords = ["kernel", "weight", "bias", "weights", "w", "b", "conv", "fc"] if any(keyword in node.name.lower() for keyword in param_keywords): # 解析该节点张量的形状 tensor_shape = node.attr["value"].tensor.tensor_shape # 计算当前参数节点的总数量(各维度大小相乘) param_num = 1 for dim in tensor_shape.dim: param_num *= dim.size total_params += param_num # 可选:打印每个参数节点的信息,方便核对 print(f"参数节点 {node.name}: {param_num} 个参数") print(f"\n.pb模型总参数数量: {total_params}") return total_params # 使用示例 if __name__ == "__main__": model_graph = load_graph("your_model_file.pb") count_pb_model_params(model_graph)
关键细节说明
节点筛选逻辑:
- 先锁定
op为Const的节点,这是冻结图参数的唯一载体 - 用关键词匹配节点名是因为不同模型的参数命名差异很大,你可以根据自己模型的实际情况修改
param_keywords列表(比如你的模型参数都叫"weight_xxx",就把"weight"加进去)
- 先锁定
调试小技巧:
如果不确定哪些节点是参数,可以先打印所有节点的名字,手动确认:print([node.name for node in graph.as_graph_def().node])然后把对应参数节点的关键词补充到筛选列表里就行。
关于参数类型:
这个方法统计的是.pb图里所有的参数(包括卷积核、全连接层权重、偏置等),因为冻结图里没有"可训练/不可训练"的区分了,所有参数都是固定的常量。
内容的提问来源于stack exchange,提问作者Yanjun
相关产品推荐
相关产品推荐

