PyTorch get all layers of model(PyTorch 获取模型的所有层)
本文介绍了PyTorch 获取模型的所有层的处理方法,对大家解决问题具有一定的参考价值,需要的朋友们下面随着小编来一起学习吧!
问题描述
在没有任何 nn.Sequence 分组的情况下,采用 pytorch 模型并获取所有层的列表的最简单方法是什么?例如,有更好的方法来做到这一点吗?
导入预训练模型定义解包模型(模型):对于我在儿童中(模型):if isinstance(i, nn.Sequential): unwrap_model(i)其他: l.append(i)模型 = pretrainedmodels.__dict__['xception'](num_classes=1000, pretrained='imagenet')l = []unwrap_model(模型)打印(升)
解决方案您可以使用
<预><代码>>>>模型 = nn.Sequential(nn.Linear(2, 2),nn.ReLU(),nn.Sequential(nn.Linear(2, 1),nn.Sigmoid()))>>>l = [model.modules() 中的模块的模块,如果不是 isinstance(module, nn.Sequential)]>>>升[线性(输入特征=2,输出特征=2,偏差=真),ReLU(),线性(输入特征=2,输出特征=1,偏差=真),Sigmoid()]modules()
方法.这是一个简单的例子:
What's the easiest way to take a pytorch model and get a list of all the layers without any nn.Sequence
groupings? For example, a better way to do this?
import pretrainedmodels
def unwrap_model(model):
for i in children(model):
if isinstance(i, nn.Sequential): unwrap_model(i)
else: l.append(i)
model = pretrainedmodels.__dict__['xception'](num_classes=1000, pretrained='imagenet')
l = []
unwrap_model(model)
print(l)
解决方案
You can iterate over all modules of a model (including those inside each Sequential
) with the modules()
method. Here's a simple example:
>>> model = nn.Sequential(nn.Linear(2, 2),
nn.ReLU(),
nn.Sequential(nn.Linear(2, 1),
nn.Sigmoid()))
>>> l = [module for module in model.modules() if not isinstance(module, nn.Sequential)]
>>> l
[Linear(in_features=2, out_features=2, bias=True),
ReLU(),
Linear(in_features=2, out_features=1, bias=True),
Sigmoid()]
这篇关于PyTorch 获取模型的所有层的文章就介绍到这了,希望我们推荐的答案对大家有所帮助,也希望大家多多支持编程学习网!
沃梦达教程
本文标题为:PyTorch 获取模型的所有层


基础教程推荐
猜你喜欢
- Dask.array.套用_沿_轴:由于额外的元素([1]),使用dask.array的每一行作为另一个函数的输入失败 2022-01-01
- Python kivy 入口点 inflateRest2 无法定位 libpng16-16.dll 2022-01-01
- 线程时出现 msgbox 错误,GUI 块 2022-01-01
- 在 Python 中,如果我在一个“with"中返回.块,文件还会关闭吗? 2022-01-01
- 使用PyInstaller后在Windows中打开可执行文件时出错 2022-01-01
- 如何在海运重新绘制中自定义标题和y标签 2022-01-01
- 如何让 python 脚本监听来自另一个脚本的输入 2022-01-01
- 筛选NumPy数组 2022-01-01
- 用于分类数据的跳跃记号标签 2022-01-01
- 何时使用 os.name、sys.platform 或 platform.system? 2022-01-01