网络编程
位置:首页>> 网络编程>> Python编程>> pytorch 实现打印模型的参数值

pytorch 实现打印模型的参数值

作者:布丁的自我修养  发布时间:2022-11-11 22:22:28 

标签:pytorch,打印模型,参数值

对于简单的网络

例如全连接层Linear

可以使用以下方法打印linear层:


fc = nn.Linear(3, 5)
params = list(fc.named_parameters())
print(params.__len__())
print(params[0])
print(params[1])

输出如下:

pytorch 实现打印模型的参数值

由于Linear默认是偏置bias的,所有参数列表的长度是2。第一个存的是全连接矩阵,第二个存的是偏置。

对于稍微复杂的网络

例如MLP


mlp = nn.Sequential(
     nn.Dropout(p=0.3),
     nn.Linear(1024, 256),
     nn.Linear(256, 64),
     nn.Linear(64, 16),
     nn.Linear(16, 1)
   )
params = list(mlp.named_parameters())
print(params.__len__())

print(params[0])
print(params[1])

print(params[2])
print(params[3])

输出:

pytorch 实现打印模型的参数值

pytorch 实现打印模型的参数值

可以发现,堆叠起来的网络,参数是依次放置的。先是全连接的权重,然后偏置。然后是下一层网络的权重+偏置。依次进行下去。

这里有4层fc,4*2=8.所以一共有8个参数矩阵。

来源:https://blog.csdn.net/budding0828/article/details/102646442

0
投稿

猜你喜欢

手机版 网络编程 asp之家 www.aspxhome.com