TensorFlow自己定义EarlyStop回调函数通过监测loss指标

简介: TensorFlow训练模型需要经过多个epoch,但是并不是epoch越多越好,很有可能训练一半的epoch时,模型的效果开始下降,这是我们需要停止训练,及时的保存模型,为了完成这种需求我们可以自定义回调函数,自动检测模型的损失,只要达到一定阈值我们手动让模型停止训练

TensorFlow训练模型需要经过多个epoch,但是并不是epoch越多越好,很有可能训练一半的epoch时,模型的效果开始下降,这是我们需要停止训练,及时的保存模型,为了完成这种需求我们可以自定义回调函数,自动检测模型的损失,只要达到一定阈值我们手动让模型停止训练


完整代码


"""

* Created with PyCharm

* 作者: 阿光

* 日期: 2022/1/4

* 时间: 10:32

* 描述:

"""

import numpy as np

import tensorflow as tf

from keras import Model

from tensorflow import keras

from tensorflow.keras.layers import *



def get_model():

   inputs = Input(shape=(784,))

   outputs = Dense(1)(inputs)

   model = Model(inputs, outputs)

   model.compile(

       optimizer=keras.optimizers.RMSprop(learning_rate=0.1),

       loss='mean_squared_error',

       metrics=['mean_absolute_error']

   )

   return model



(x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data()

x_train = x_train.reshape(-1, 784).astype('float32') / 255.0

x_test = x_test.reshape(-1, 784).astype('float32') / 255.0


x_train = x_train[:1000]

y_train = y_train[:1000]

x_test = x_test[:1000]

y_test = y_test[:1000]



class CustomEarlyStoppingAtMinLoss(keras.callbacks.Callback):

   def __init__(self, patience=0):

       super(CustomEarlyStoppingAtMinLoss, self).__init__()

       self.patience = patience

       self.best_weights = None

       self.wait = 0

       self.stopped_epoch = 0

       self.best = np.Inf


   def on_train_begin(self, logs=None):

       pass


   def on_epoch_end(self, epoch, logs=None):

       current = logs.get("loss")

       if np.less(current, self.best):

           self.best = current

           self.wait = 0

           self.best_weights = self.model.get_weights()

       else:

           self.wait += 1

           if self.wait >= self.patience:

               self.stopped_epoch = epoch

               self.model.stop_training = True

               print("Restoring model weights from the end of the best epoch.")

               self.model.set_weights(self.best_weights)


   def on_train_end(self, logs=None):

       if self.stopped_epoch > 0:

           print("Epoch %05d: early stopping" % (self.stopped_epoch + 1))



model = get_model()

model.fit(

   x_train,

   y_train,

   batch_size=128,

   epochs=10,

   verbose=1,

   validation_split=0.5,

   callbacks=[CustomEarlyStoppingAtMinLoss()],

)

目录
相关文章
|
NoSQL Linux 程序员
Linux:gdb调试器的解析+使用(超详细版)
Linux:gdb调试器的解析+使用(超详细版)
854 1
|
网络协议 测试技术 Linux
中国移动ML302模组(4G Cat.1 通信模组)TencentOS-tiny AT模组框架适配
中国移动ML302模组(4G Cat.1 通信模组)TencentOS-tiny AT模组框架适配
733 0
|
10月前
|
存储 安全 数据安全/隐私保护
windows远程桌面配置CA证书
本文介绍如何在Windows系统中导入TLS证书并配置其权限与应用。通过MMC控制台添加证书管理单元,导入PFX格式证书,设置私钥访问权限,并使用WMIC命令将证书指纹绑定至远程桌面服务,实现安全加密连接。
1008 6
|
机器学习/深度学习 监控 算法
基于mediapipe深度学习的手势数字识别系统python源码
本内容涵盖手势识别算法的相关资料,包括:1. 算法运行效果预览(无水印完整程序);2. 软件版本与配置环境说明,提供Python运行环境安装步骤;3. 部分核心代码,完整版含中文注释及操作视频;4. 算法理论概述,详解Mediapipe框架在手势识别中的应用。Mediapipe采用模块化设计,包含Calculator Graph、Packet和Subgraph等核心组件,支持实时处理任务,广泛应用于虚拟现实、智能监控等领域。
|
数据采集 搜索推荐 项目管理
通用型埋点系统完整开源方案-ClkLog新升级更强大、更易用
我们希望ClkLog开源社区版,不是“精简试用版”,而是一个真正能被部署和使用的完整方案。 过去这一年,我们一直在倾听大家的反馈,并不断思考:一款开源行为分析系统,真正顺利地被用起来,需要具备哪些要素和功能? 为了让大家在使用过程中更流畅更便捷,ClkLog开源社区版迎来了一次新升级! 现在上Gitee、Github、GitCode 即可获取最新的更新代码
|
算法 数据可视化 数据挖掘
【数据挖掘】密度聚类DBSCAN讲解及实战应用(图文解释 附源码)
【数据挖掘】密度聚类DBSCAN讲解及实战应用(图文解释 附源码)
1686 1
|
10月前
|
缓存 JSON 搜索推荐
拼多多商品详情API接口指南
拼多多商品详情API是开放平台提供的商品数据查询接口,支持获取商品信息、价格、库存、销量、评价及促销等关键数据,返回结构化JSON格式。适用于电商数据分析、价格监测、竞品分析与个性化推荐场景,配合缓存、批量请求与签名优化策略,提升调用效率与系统稳定性。(238字)
1217 1
|
索引
【Qt 学习笔记】Qt常用控件 | 多元素控件 | List Widget的说明及介绍
【Qt 学习笔记】Qt常用控件 | 多元素控件 | List Widget的说明及介绍
1736 3
|
JavaScript 前端开发 Java
如何使用正则表达式来匹配电子邮件地址?
如何使用正则表达式来匹配电子邮件地址?
1468 3

热门文章

最新文章