mxnet下如何查看中间结果
https://blog.csdn.net/disen10/article/details/79376631
固定权重:https://www.cnblogs.com/chenyliang/p/6780019.html
固定权重:https://discuss.gluon.ai/t/topic/1164
查看权重
在训练过程中,有时候我们为了debug而需要查看中间某一步的权重信息,在mxnet中,我们可以很方便的调用get_params()方法来得到权重信息。
- '''
- 查看权重示例代码
- 转载时注明地址:http://blog.csdn.net/u010414386?viewmode=contents
- '''
- import mxnet as mx
- sym, arg_params, aux_params = mx.model.load_checkpoint('resnet-50',0)#载入模型
- mod = mx.mod.Module(symbol=sym,context=mx.gpu()) #创建Module
- mod.bind(for_training=False,data_shapes=[('data',(1,3,224,224))]) #绑定,此代码为预测代码,所以training参数设为False
- mod.set_params(arg_params,aux_params)
- import numpy as np
- import cv2
- def get_image(filename):
- img = cv2.imread(filename)
- img = cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
- img = cv2.resize(img,(224,224))
- img = np.swapaxes(img,0,2)
- img = np.swapaxes(img,1,2)
- img = img[np.newaxis,:]
- return img
- from collections import namedtuple
- Batch = namedtuple('Batch',['data'])
- img = get_image('val_1000/0.jpg') #获取图片
- mod.forward(Batch([mx.nd.array(img)])) #预测结果
- ################################################
- #debug模式下,获取权重信息
- keys = mod.get_params()[0].keys() # 列出所有权重名称
- conv_w = mod.get_params()[0]['conv0_weight'] #获取想要查看的权重信息,如conv_weight
- print conv_w.asnumpy() #查看具体数值
- ################################################
- prob = mod.get_outputs()[0].asnumpy()
- y = np.argsort(np.squeeze(prob))[::-1]
- print('truth label %d; top-1 predict label %d' % (val_label[0], y[0]))
- 1
- 2
- 3
- 4
- 5
- 6
- 7
- 8
- 9
- 10
- 11
- 12
- 13
- 14
- 15
- 16
- 17
- 18
- 19
- 20
- 21
- 22
- 23
- 24
- 25
- 26
- 27
- 28
- 29
- 30
- 31
- 32
- 33
查看中间输出结果
由于mxnet的网络由symbol组成,而symbol又属于符号式编程,所以我们不能像上面查看权重一样直接查看,我们需要把我们想看的输出结果保存下来。
- '''
- 方法一
- 查看中间结果代码
- 转载时注明地址:http://blog.csdn.net/u010414386?viewmode=contents
- '''
- import mxnet as mx
- net = mx.symbol.Variable('data')
- fc1 = mx.symbol.FullyConnected(data=net, name='fc1', num_hidden=128)
- net = mx.symbol.Activation(data=fc1, name='relu1', act_type="relu")
- net = mx.symbol.FullyConnected(data=net, name='fc2', num_hidden=64)
- out = mx.symbol.SoftmaxOutput(data=net, name='softmax')
- # 通过把两个输出组成一个group来得到自己需要查看的中间层输出结果
- group = mx.symbol.Group([fc1, out])
- print group.list_outputs()
- 1
- 2
- 3
- 4
- 5
- 6
- 7
- 8
- 9
- 10
- 11
- 12
- 13
- 14
- '''
- 方法二
- 有时候我们使用别人的模型,所以无法像方法一一样在定义模型的时候就确定需要查看的中间层输出结果,
- 这时候我们使用get_internals()方法来查找自己需要查看的中间层
- 转载时注明地址:http://blog.csdn.net/u010414386?viewmode=contents
- '''
- import mxnet as mx
- sym, arg_params, aux_params = mx.model.load_checkpoint('resnet-50',0)#载入模型
- ########################################################################
- args = sym.get_internals().list_outputs() #获得所有中间输出
- internals = model.symbol.get_internals()
- fc1 = internals['fc1_output']
- conv = internals['stage4_unit3_conv1_output']
- group = mx.symbol.Group([fc1, sym, conv]) #把需要输出的结果按group方式组合起来,这样就可以得到中间层的输出
- #########################################################################
- mod = mx.mod.Module(symbol=group,context=mx.gpu()) #创建Module
- mod.bind(for_training=False,data_shapes=[('data',(1,3,224,224))]) #绑定,此代码为预测代码,所以training参数设为False
- mod.set_params(arg_params,aux_params)
- import numpy as np
- import cv2
- def get_image(filename):
- img = cv2.imread(filename)
- img = cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
- img = cv2.resize(img,(224,224))
- img = np.swapaxes(img,0,2)
- img = np.swapaxes(img,1,2)
- img = img[np.newaxis,:]
- return img
- from collections import namedtuple
- Batch = namedtuple('Batch',['data'])
- img = get_image('val_1000/0.jpg') #获取图片
- mod.forward(Batch([mx.nd.array(img)])) #预测结果
- prob = mod.get_outputs()[0].asnumpy()
- y = np.argsort(np.squeeze(prob))[::-1]
- print('truth label %d; top-1 predict label %d' % (val_label[0], y[0]))
mxnet下如何查看中间结果的更多相关文章
- Linux下如何查看版本信息
Linux下如何查看版本信息, 包括位数.版本信息以及CPU内核信息.CPU具体型号等等,整个CPU信息一目了然. 1.# uname -a (Linux查看版本当前操作系统内核信息) L ...
- Linux下怎么查看当前系统的版本
Linux下怎么查看当前系统的版本: uname -r 功能说明:uname用来获取电脑和操作系统的相关信息. 语 法:uname [-amnrsvpio][--help][--version] ...
- 在windows和linux下如何查看80端口占用情况?是被哪个进程占用?如何终止等
一.在windows下如何查看80端口占用情况?是被哪个进程占用?如何终止等 这里主要是用到windows下的DOS工具,点击"开始"--"运行",输入&quo ...
- 在linux下,查看一个运行中的程序, 占用了多少内存
1. 在linux下,查看一个运行中的程序, 占用了多少内存, 一般的命令有 (1). ps aux: 其中 VSZ(或VSS)列 表示,程序占用了多少虚拟内存. RSS列 表示, 程序占用了多少物 ...
- linux下如何查看mysql、apache是否安装,并卸载
--linux下如何查看mysql.apache是否安装,并卸载? http://blog.163.com/dengxiuhua126@126/blog/static/1186077720137311 ...
- Linux 下实时查看日志
Linux 下实时查看日志 cat /var/log/*.log 如果日志在更新,如何实时查看 tail -f /var/log/messages 还可以使用 watch -d -n 1 cat /v ...
- Linux下如何查看tomcat是否安装、启动、文件路径、进程ID
Linux下如何查看tomcat是否安装.启动.文件路径.进程ID 在Linux系统下,Tomcat使用命令的操作! 检测是否有安装了Tomcat: rpm -qa|grep tomcat 查看Tom ...
- Linux下内存查看命令
在Linux下面,我们常用top命令来查看系统进程,top也能显示系统内存.我们常用的Linux下查看内容的专用工具是free命令. Linux下内存查看命令free详解: 在Linux下查看内存我们 ...
- Linux之Ubuntu下如何查看已安装的软件/库文件【摘抄】
本文属于实用性质,且属于摘抄别处,出自:[Ubuntu 下如何查看已安装的软件](http://blog.csdn.net/m1205979825/article/details/40855583) ...
随机推荐
- PHP做APP接口时,如何保证接口的安全性??????????
PHP做APP接口时,如何保证接口的安全性? 1.当用户登录APP时,使用https协议调用后台相关接口,服务器端根据用户名和密码时生成一个access_key,并将access_key保存在sess ...
- Java学习路径及练手项目合集
Java 在编程语言排行榜中一直位列前排,可知 Java 语言的受欢迎程度了. 实验楼上的[Java 学习路径]中将首先完成 Java基础.JDK.JDBC.正则表达式等基础实验,然后进阶到 J2SE ...
- React.createClass 、React.createElement、Component
react里面有几个需要区别开的函数 React.createClass .React.createElement.Component 首选看一下在浏览器的下面写法: <div id=" ...
- 如何设置locale
什么是 locale? 是根据计算机用户所使用的语言,所在国家或者地区,以及当地的文化传统所定义的一个软件运行时的语言环境 locale定义文件放在目录 /usr/share/i18n/locales ...
- fullpage插件在移动端弹出键盘页面特殊处理
fullpage插件大家都很熟悉 jquery一款全屏上下滑动的插件. 最近做公司一个活动移动端使用fullpage插件填写input的时候遇见一个问题,手机自带的键盘弹出的时候会把页面顶出去,页面错 ...
- iOS 设计模式-NSNotificationCenter 通知中心
通知介绍 每一个应用程序都有一个通知中心(NSNotificationCenter)实例,专门负责协助不同对象之间的消息通信 任何一个对象都可以向通知中心发布通知(NSNotification),描述 ...
- 《linux就该这么学》找到一本不错的Linux电子书,《Linux就该这么学》。
本帖不是广告贴,只是感觉有好的工具书而已 本书是由全国多名红帽架构师(RHCA)基于最新Linux系统共同编写的高质量Linux技术自学教程,极其适合用于Linux技术入门教程或讲课辅助教材,目前是国 ...
- vue框架(三)_vue引入jquery、bootstrap
一.vue安装jquery 1.按照之前博客的内容,新建一个vue工程. 2.在项目文件夹下,使用命令npm install jquery --save-dev 引入jquery. 3.在build/ ...
- uvalive 3887 Slim Span
题意: 一棵生成树的苗条度被定义为最长边与最小边的差. 给出一个图,求其中生成树的最小苗条度. 思路: 最开始想用二分,始终想不到二分终止的条件,所以尝试暴力枚举最小边的长度,然后就AC了. 粗略估计 ...
- 定时调度任务quartz
依赖 <!-- quartz --> <dependency> <groupId>org.quartz-scheduler</groupId> < ...