PyTorch加载.pt文件:如何打印模型参数形状与参数值
处理PyTorch MLP模型参数字典的方法
1. 加载模型参数字典
首先用torch.load()加载.pt文件,得到参数字典:
import torch # 替换为你的模型文件路径 model_dict = torch.load("your_model.pt")
2. 打印所有参数的形状
遍历字典的键值对,就能看到每个参数的名称和对应形状:
for name, tensor in model_dict.items(): print(f"{name}: {tensor.shape}")
3. 获取输入层到隐藏层的权重
PyTorch中全连接层(nn.Linear)的权重形状是[输出维度, 输入维度],所以输入层到隐藏层的权重原始形状会是[32, 168],转置后就是你需要的168*32矩阵:
# 根据第二步打印的参数名称,找到输入到隐藏层的权重键名(通常类似fc1.weight、layers.0.weight) input_hidden_weight = model_dict["fc1.weight"] # 转置得到168x32的矩阵 target_weight = input_hidden_weight.T print("输入层到隐藏层的权重矩阵(168x32):") print(target_weight)
注意:如果你的模型参数命名不同,比如用了Sequential容器,键名可能是
0.weight,以第二步打印的实际名称为准。
内容的提问来源于stack exchange,提问作者Leon.L
相关产品推荐
相关产品推荐

