如何判断TensorFlow/Keras生成的Protocol Buffer(.pb)文件是否相同?
解决TensorFlow .pb文件一致性校验的问题
我之前也踩过这个坑!直接用hashlib.md5对整个.pb文件哈希确实行不通——哪怕是用同一脚本冻结同一模型,生成的文件哈希也可能不一样。这是因为TensorFlow在冻结图的过程中,会带入一些非模型核心逻辑的元数据,比如节点的生成顺序、会话的临时状态、甚至是某些内部辅助节点的随机命名,这些都会改变文件的二进制内容,但模型的结构和权重其实是完全一致的。
下面分享几个可行的解决方案,按可靠性排序:
方案1:提取模型结构+权重生成唯一指纹(最可靠)
这个方法会忽略所有无关元数据,只提取模型的核心信息——节点结构、输入输出关系、权重值——然后将这些信息序列化后计算哈希。不管.pb文件的其他内容怎么变,只要模型的结构和权重一致,指纹就会相同。
代码实现
import tensorflow as tf import hashlib import json def get_model_fingerprint(pb_path): # 加载.pb文件到GraphDef with tf.io.gfile.GFile(pb_path, 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) fingerprint_data = [] # 按节点名称排序,避免节点顺序不同导致指纹差异 for node in sorted(graph_def.node, key=lambda x: x.name): node_info = { 'name': node.name, 'op': node.op, 'inputs': list(node.input), 'attrs': {} } # 只保留影响模型逻辑的关键属性,忽略device等无关项 for attr_name, attr_value in node.attr.items(): if attr_name in ['dtype', 'shape', 'value']: if attr_name == 'value': try: # 将Const节点的权重张量转换成可序列化的列表 tensor = tf.make_ndarray(attr_value.tensor) node_info['attrs'][attr_name] = tensor.tolist() except: # 非张量类型的value,保留字符串表示 node_info['attrs'][attr_name] = str(attr_value) else: # dtype和shape转成字符串格式 node_info['attrs'][attr_name] = str(attr_value) fingerprint_data.append(node_info) # 序列化数据时强制排序键值对,保证输出一致 serialized = json.dumps(fingerprint_data, sort_keys=True).encode('utf-8') # 计算MD5指纹 return hashlib.md5(serialized).hexdigest() # 使用示例 fp1 = get_model_fingerprint("my_model1.pb") fp2 = get_model_fingerprint("my_model2.pb") print(f"模型1指纹:{fp1}") print(f"模型2指纹:{fp2}") print(f"模型是否完全相同:{fp1 == fp2}")
方案2:仅校验权重哈希(快速但不全面)
如果只关心模型的权重是否一致,可以单独提取所有Const节点的权重值,计算它们的组合哈希。这个方法更快,但无法检测结构差异(比如两个模型权重相同但网络结构不同的情况)。
代码实现
import tensorflow as tf import hashlib def get_weights_hash(pb_path): graph = tf.Graph() with graph.as_default(): # 加载.pb文件到图中 with tf.io.gfile.GFile(pb_path, 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) tf.import_graph_def(graph_def, name='') weights_hash = hashlib.md5() # 遍历所有Const节点,提取权重并更新哈希 with tf.compat.v1.Session(graph=graph) as sess: for node in graph_def.node: if node.op == 'Const': tensor = sess.run(f"{node.name}:0") # 将张量转换成字节串更新哈希 weights_hash.update(tensor.tobytes()) return weights_hash.hexdigest()
方案3:优化冻结脚本减少差异(治标不治本)
如果你想尽可能让生成的.pb文件哈希一致,可以修改你的冻结代码,固定一些可能导致无序的变量:
- 把
set转换成sorted,避免集合的无序性导致节点顺序变化 - 对输出名称和冻结变量名称进行排序
修改后的冻结代码片段:
# 原代码中set是无序的,改成排序后的列表 freeze_var_names = sorted(list(set(v.op.name for v in tf.global_variables()).difference(keep_var_names or []))) # 对输出名称排序,保证顺序固定 output_names = sorted(output_names or []) output_names += sorted([v.op.name for v in tf.global_variables()])
不过这个方法只能减少差异,无法完全避免——TensorFlow内部仍可能有一些不可控的元数据写入,所以还是推荐方案1作为最终的校验方法。
内容的提问来源于stack exchange,提问作者jss367
相关产品推荐
相关产品推荐

