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

TensorFlow中与torch.load()等效的功能是什么?如何查看模型参数?

TensorFlow中等效于torch.load查看模型参数的实现

和PyTorch的torch.load('/filepath')功能对应,TensorFlow根据你使用的模型存储格式,有以下几种常用实现:

  • SavedModel格式(TensorFlow 2.x 默认导出格式)
    直接调用tf.keras.models.load_model加载完整模型后即可查看所有参数:

    import tensorflow as tf
    
    # 加载模型,对应torch.load的基础功能
    model = tf.keras.models.load_model('/你的模型存储路径')
    
    # 查看所有可训练参数(权重、偏置等)
    trainable_params = model.trainable_weights
    for param in trainable_params:
        print(f"参数名: {param.name}, 形状: {param.shape}")
        # 输出参数具体数值
        print(param.numpy())
    
    # 查看所有参数,包含不可训练参数(如BatchNorm层的滑动均值、方差)
    all_params = model.get_weights()
    

    如果你不需要加载完整模型结构,只需要读取checkpoint里的参数,可以用tf.train.load_checkpoint:

    # 初始化checkpoint读取器
    ckpt_reader = tf.train.load_checkpoint('/你的checkpoint文件路径')
    # 输出所有参数的名称和对应形状
    print(ckpt_reader.get_variable_to_shape_map())
    # 读取指定参数的数值
    param = ckpt_reader.get_tensor('目标参数名')
    
  • HDF5格式(.h5/.hdf5后缀,Keras常用旧存储格式)
    同样使用load_model接口加载后操作即可,和SavedModel格式的参数查看逻辑完全一致:

    model = tf.keras.models.load_model('model_weights.h5')
    print(model.get_weights())
    

如果你使用的是TensorFlow 1.x版本通过tf.train.Saver保存的checkpoint文件,也可以直接用上面提到的tf.train.load_checkpoint接口读取参数,不需要切换到TF1.x的会话运行模式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 19:15:08