广告
返回顶部
首页 > 资讯 > 精选 >怎么在Pytorch模型中将读取的pth文件参数转换成numpy矩阵
  • 224
分享到

怎么在Pytorch模型中将读取的pth文件参数转换成numpy矩阵

2023-06-06 18:06:56 224人浏览 独家记忆
摘要

怎么在PyTorch模型中将读取的pth文件参数转换成numpy矩阵?很多新手对此不是很清楚,为了帮助大家解决这个难题,下面小编将为大家详细讲解,有这方面需求的人可以来学习下,希望你能有所收获。Pytorch给了很方便的读取参数接口:nn.

怎么在PyTorch模型中将读取的pth文件参数转换成numpy矩阵?很多新手对此不是很清楚,为了帮助大家解决这个难题,下面小编将为大家详细讲解,有这方面需求的人可以来学习下,希望你能有所收获。

Pytorch给了很方便的读取参数接口:

nn.Module.parameters()

直接看demo:

from torchvision.models.alexnet import alexnet model = alexnet(pretrained=True).eval().cuda()parameters = model.parameters()for p in parameters:  numpy_para = p.detach().cpu().numpy()  print(type(numpy_para))  print(numpy_para.shape)

上面得到的numpy_para就是numpy参数了~

Note:

model.parameters()是以一个生成器的形式迭代返回每一层的参数。所以用for循环读取到各层的参数,循环次数就表示层数。

而每一层的参数都是torch.nn.parameter.Parameter类型,是Tensor的子类,所以直接用tensor转numpy(即p.detach().cpu().numpy())的方法就可以直接转成numpy矩阵。

方便又好用,爆赞~

补充:pytorch训练好的.pth模型转换为.pt

python训练好的.pth文件转为.pt

import torchimport torchvisionfrom unet import UNetmodel = UNet(3, 2)#自己定义的网络模型model.load_state_dict(torch.load("best_weights.pth"))#保存的训练模型model.eval()#切换到eval()example = torch.rand(1, 3, 320, 480)#生成一个随机输入维度的输入traced_script_module = torch.jit.trace(model, example)traced_script_module.save("model.pt")

看完上述内容是否对您有帮助呢?如果还想对相关知识有进一步的了解或阅读更多相关文章,请关注编程网精选频道,感谢您对编程网的支持。

--结束END--

本文标题: 怎么在Pytorch模型中将读取的pth文件参数转换成numpy矩阵

本文链接: https://www.lsjlt.com/news/248194.html(转载时请注明来源链接)

有问题或投稿请发送至: 邮箱/279061341@qq.com    QQ/279061341

本篇文章演示代码以及资料文档资料下载

下载Word文档到电脑,方便收藏和打印~

下载Word文档
猜你喜欢
  • 怎么在Pytorch模型中将读取的pth文件参数转换成numpy矩阵
    怎么在Pytorch模型中将读取的pth文件参数转换成numpy矩阵?很多新手对此不是很清楚,为了帮助大家解决这个难题,下面小编将为大家详细讲解,有这方面需求的人可以来学习下,希望你能有所收获。Pytorch给了很方便的读取参数接口:nn....
    99+
    2023-06-06
软考高级职称资格查询
编程网,编程工程师的家园,是目前国内优秀的开源技术社区之一,形成了由开源软件库、代码分享、资讯、协作翻译、讨论区和博客等几大频道内容,为IT开发者提供了一个发现、使用、并交流开源技术的平台。
  • 官方手机版

  • 微信公众号

  • 商务合作