使用回调函数及tensorboard实现网络训练实时监控

2024-04-30 22:08

本文主要是介绍使用回调函数及tensorboard实现网络训练实时监控,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

神经网络开发的一大特点是, 一旦我们把大规模数据输入网络进行分析时,你的感觉就像抛出一只纸飞机,除了抛出那一刻你拥有控制力外,一旦离手,它怎么飞怎么飘就不再是你能控制得了。神经网络代码的运行就有这个特点,我们不能像平常程序那样设置断点,然后单步调试,一旦运行后,我们只能观察结果。令人郁闷的是,很多时候训练非常耗时,你跑完几个小时后突然发现代码中存在bug,于是你停下程序,修正后你又得等待好几个小时。

幸运的是,keras框架早就意识到这一点,它提供了相应机制能让我们随时监控网络的运行状况。通过前面章节我们看到,通常情况下我们不知道需要几个循环,网络才能达到最佳效果,我们往往让网络训练很多个循环,直到出现过度拟合时,我再观察训练过程数据,从中找到网络达到最佳状况所需的训练循环,然后我们重新设置循环次数后,再将网络重头跑一遍,这是非常耗时,效率低下的工作。

一个好的解决办法是提供一种监控机制,一旦发现网络对校验数据的判断准确率没有明显提升后就停止训练。keras提供了回调机制让我们随时监控网络的训练状况。当我们只需fit函数启动网络训练时,我们可以提供一个回调对象,网络每训练完一个流程后,它会回调我们提供的函数,在函数里我们可以访问网络所有参数从而知道网络当前运行状态,此时我们可以采取多种措施,例如终止训练流程,保存网络所有参数,加载新参数等,甚至我们能改变网络的运行状态。

keras提供的回调具体来说可以让我们完成几种操作,一种是存储网络当前所有参数;一种是停止训练流程;一种是调节与训练相关的某些参数,例如学习率,一种是输出网络状态信息,或者对网络内部状况进行视觉化输出,我们看一些代码例子:

import keras
callbacks_list = [#停止训练流程,一旦网络对校验数据的判断率不再提升,patience表示在两次循环间判断率没改进时就停止keras.callbacks.EarlyStopping(monitor='acc', patience=1),'''在每次训练循环结束时将当前参数存入文件my_model.h5,后两个参数表明当网络判断率没有提升时,不存储参数'''keras.callbacks.ModelCheckPoint(filepat='my_model.h5',monitor='val_loss',save_best_only=True),
'''如果网络对校验数据的判断率在10次训练循环内一直没有提升,下面回调将修改学习率'''keras.callbacks.ReduceLROnPlateau(monitor='val_loss',factor=0.1,patience=10,)
]model.compile(optimizer='rmsprop',loss='binary_crossentropy',metrics=['acc'])
'''
由于回调函数中会监控网络对校验数据判断的准确率,因此训练网络时必须传入校验数据
'''
model.fit(x, y, epochs = 10, callbacks = callbacks_list,validation_data = (x_val, y_val))

要想训练出一个精准的网络,一个重要前提是我们能时刻把握网络内部状态的变化情况,如果这些变化能够以视觉化的方式实时显示出来,那么我们就能方便的掌握网络内部的状态变化,keras框架附带的一个组件叫tensorboard能有效的帮我们实现这点,接下来我们构造一个网络,然后输入数据训练网络,然后激活tensorboard,通过可视化的方式看看网络在训练过程中的变化:

import keras;
from keras import layers
from keras.datasets import imdb
from keras.preprocessing import sequencemax_features = 2000
max_len = 500(x_train, y_train), (x_test, y_test) = imdb.load_data(num_words = max_features)
x_train = sequence.pad_sequences(x_train, maxlen=max_len)
x_test = sequence.pad_sequence(x_test, maxlen = max_len)model = keras.models.Sequential()
model.add(layers.Embedding(max_features, 128, input_length = max_len,name = 'embed'))
model.add(layers.Conv1D(32, 7, activation='relu'))
model.add(layers.MaxPooling1D(5))
model.add(layers.Conv1D(32, 7, activation='relu'))
model.add(layers.GlobalMaxPooling1D())
model.add(layers.Dense(1))
model.summary()
model.compile(optimizer = 'rmsprop', loss = 'binary_crossentropy',metrics = ['acc'])

上面代码我们以前讲解过,这里的重点不再是理解它的逻辑,而是让它跑起来,然后我们使用tensorboard观察网络内在状态的变化,要使用tensorboard,我们需要创建一个目录用于存储它运行时生成的日志:

!mkdir my_log_dir

接着我们给网络注入一个回调钩子,让它在运行时把内部信息传递给tensorbaord组件:

callbacks = [keras.callbacks.TensorBoard(log_dir='my_log_dir',#每隔一个训练循环就用柱状图显示信息histogram_freq = 1,embeddings_freq = 1)
]history = model.fit(x_train, y_train,epochs = 20,batch_size = 128,validation_split = 0.2,callbacks = callbacks)

执行上面代码启动训练后,我们在控制台输入如下命令:

conda activate tensorflow
tensorboard --log_dir=my_log_dir

第一句命令用于激活安装了tensorflow的环境,第二句启动tensorbaord服务器。此时在浏览器里输入:http://localhost:6006就可以打开可视化环境,如下图:

屏幕快照 2019-01-08 下午4.44.36.png

点击histogram,我们可以看到网络内部状态变化以柱状图的方式展现出来:

屏幕快照 2019-01-08 下午4.46.20.png

更强大的是,它会把我们训练的单词向量以可视化的方式展现出来,点击Projector,你会看到如下三维动画:

屏幕快照 2019-01-08 下午4.49.10.png

它使用t-SNE可视化算法把高维向量转换到二维空间上进行展示。点击Graph按钮,它会把网络的模型图绘制出来,让你了解网络的层次结构:

屏幕快照 2019-01-08 下午4.52.27.png

有了回调函数和tensorboard组件的帮助,我们不用再将网络看做是一个无法窥探的黑盒子,通过tensorboard,我们可以在非常详实的视觉辅助下掌握网络的训练流程以及内部状态变化。

更多技术信息,包括操作系统,编译器,面试算法,机器学习,人工智能,请关照我的公众号:
这里写图片描述

更多内容,请点击进入csdn学院

这篇关于使用回调函数及tensorboard实现网络训练实时监控的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/950042

相关文章

PostgreSQL中rank()窗口函数实用指南与示例

《PostgreSQL中rank()窗口函数实用指南与示例》在数据分析和数据库管理中,经常需要对数据进行排名操作,PostgreSQL提供了强大的窗口函数rank(),可以方便地对结果集中的行进行排名... 目录一、rank()函数简介二、基础示例:部门内员工薪资排名示例数据排名查询三、高级应用示例1. 每

使用Python删除Excel中的行列和单元格示例详解

《使用Python删除Excel中的行列和单元格示例详解》在处理Excel数据时,删除不需要的行、列或单元格是一项常见且必要的操作,本文将使用Python脚本实现对Excel表格的高效自动化处理,感兴... 目录开发环境准备使用 python 删除 Excphpel 表格中的行删除特定行删除空白行删除含指定

全面掌握 SQL 中的 DATEDIFF函数及用法最佳实践

《全面掌握SQL中的DATEDIFF函数及用法最佳实践》本文解析DATEDIFF在不同数据库中的差异,强调其边界计算原理,探讨应用场景及陷阱,推荐根据需求选择TIMESTAMPDIFF或inte... 目录1. 核心概念:DATEDIFF 究竟在计算什么?2. 主流数据库中的 DATEDIFF 实现2.1

深入理解Go语言中二维切片的使用

《深入理解Go语言中二维切片的使用》本文深入讲解了Go语言中二维切片的概念与应用,用于表示矩阵、表格等二维数据结构,文中通过示例代码介绍的非常详细,需要的朋友们下面随着小编来一起学习学习吧... 目录引言二维切片的基本概念定义创建二维切片二维切片的操作访问元素修改元素遍历二维切片二维切片的动态调整追加行动态

Linux下删除乱码文件和目录的实现方式

《Linux下删除乱码文件和目录的实现方式》:本文主要介绍Linux下删除乱码文件和目录的实现方式,具有很好的参考价值,希望对大家有所帮助,如有错误或未考虑完全的地方,望不吝赐教... 目录linux下删除乱码文件和目录方法1方法2总结Linux下删除乱码文件和目录方法1使用ls -i命令找到文件或目录

MySQL中的LENGTH()函数用法详解与实例分析

《MySQL中的LENGTH()函数用法详解与实例分析》MySQLLENGTH()函数用于计算字符串的字节长度,区别于CHAR_LENGTH()的字符长度,适用于多字节字符集(如UTF-8)的数据验证... 目录1. LENGTH()函数的基本语法2. LENGTH()函数的返回值2.1 示例1:计算字符串

prometheus如何使用pushgateway监控网路丢包

《prometheus如何使用pushgateway监控网路丢包》:本文主要介绍prometheus如何使用pushgateway监控网路丢包问题,具有很好的参考价值,希望对大家有所帮助,如有错误... 目录监控网路丢包脚本数据图表总结监控网路丢包脚本[root@gtcq-gt-monitor-prome

SpringBoot+EasyExcel实现自定义复杂样式导入导出

《SpringBoot+EasyExcel实现自定义复杂样式导入导出》这篇文章主要为大家详细介绍了SpringBoot如何结果EasyExcel实现自定义复杂样式导入导出功能,文中的示例代码讲解详细,... 目录安装处理自定义导出复杂场景1、列不固定,动态列2、动态下拉3、自定义锁定行/列,添加密码4、合并

mybatis执行insert返回id实现详解

《mybatis执行insert返回id实现详解》MyBatis插入操作默认返回受影响行数,需通过useGeneratedKeys+keyProperty或selectKey获取主键ID,确保主键为自... 目录 两种方式获取自增 ID:1. ​​useGeneratedKeys+keyProperty(推

Spring Boot集成Druid实现数据源管理与监控的详细步骤

《SpringBoot集成Druid实现数据源管理与监控的详细步骤》本文介绍如何在SpringBoot项目中集成Druid数据库连接池,包括环境搭建、Maven依赖配置、SpringBoot配置文件... 目录1. 引言1.1 环境准备1.2 Druid介绍2. 配置Druid连接池3. 查看Druid监控