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

如何判断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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 06:49:50