pytorch 获取层权重,对特定层注入hook, 提取中间层输出的方法
2019-08-17 09:44
1671 查看
如下所示:
#获取模型权重 for k, v in model_2.state_dict().iteritems(): print("Layer {}".format(k)) print(v)
#获取模型权重 for layer in model_2.modules(): if isinstance(layer, nn.Linear): print(layer.weight)
#将一个模型权重载入另一个模型 model = VGG(make_layers(cfg['E']), **kwargs) if pretrained: load = torch.load('/home/huangqk/.torch/models/vgg19-dcbb9e9d.pth') load_state = {k: v for k, v in load.items() if k not in ['classifier.0.weight', 'classifier.0.bias', 'classifier.3.weight', 'classifier.3.bias', 'classifier.6.weight', 'classifier.6.bias']} model_state = model.state_dict() model_state.update(load_state) model.load_state_dict(model_state) return model
# 对特定层注入hook def hook_layers(model): def hook_function(module, inputs, outputs): recreate_image(inputs[0]) print(model.features._modules) first_layer = list(model.features._modules.items())[0][1] first_layer.register_forward_hook(hook_function)
#获取层 x = someinput for l in vgg.features.modules(): x = l(x) modulelist = list(vgg.features.modules()) for l in modulelist[:5]: x = l(x) keep = x for l in modulelist[5:]: x = l(x)
# 提取vgg模型的中间层输出 # coding:utf8 import torch import torch.nn as nn from torchvision.models import vgg16 from collections import namedtuple class Vgg16(torch.nn.Module): def __init__(self): super(Vgg16, self).__init__() features = list(vgg16(pretrained=True).features)[:23] # features的第3,8,15,22层分别是: relu1_2,relu2_2,relu3_3,relu4_3 self.features = nn.ModuleList(features).eval() def forward(self, x): results = [] for ii, model in enumerate(self.features): x = model(x) if ii in {3, 8, 15, 22}: results.append(x) vgg_outputs = namedtuple("VggOutputs", ['relu1_2', 'relu2_2', 'relu3_3', 'relu4_3']) return vgg_outputs(*results)
以上这篇pytorch 获取层权重,对特定层注入hook, 提取中间层输出的方法就是小编分享给大家的全部内容了,希望能给大家一个参考
您可能感兴趣的文章:
相关文章推荐
- PyTorch快速搭建神经网络及其保存提取方法详解
- c#获取两个特定字符之间的内容并输出的方法
- C#获取存储过程返回值和输出参数值的方法
- ASP.NET Core DI手动获取注入对象的方法
- Java中获取特定符号中间字符串子串的方法
- 【转载】linux c程序中获取shell脚本输出的实现方法
- http获取输入流输出流的参数及方法
- linux c程序中获取shell脚本输出的实现方法
- 普通静态类方法获取Spring注入的Been实体
- android-获取网络时间、获取特定时区时间、时间同步的方法
- php获取多维数组某个特定键(数组下标)的所有值,具体总结下其余的方法
- 两个使用正则表达式来获取字符串中特定子串的方法
- 由Spring管理的bean,不使用注入的方式来获取bean的方法
- 谈谈获取XML格式数据中特定节点值的方法
- 获取手机音频输出设备方法
- Cocoa中用NSTask执行外部命令并获取输出结果的方法
- 使用Mono Cecil 动态获取运行时数据 (Atribute形式 进行注入 用于写Log) [此文报考 xxx is declared in another module and needs to be imported的解决方法]-摘自网络
- C#获取命令行输出内容的方法
- dedecms获取图片集多张图片实现方法(循环输出)
- easyui中combotree循环获取父节点至根节点并输出路径实现方法