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

