政安晨:【Keras机器学习实践要点】(七)—— 使用TensorFlow自定义fit()

政安晨的个人主页政安晨

欢迎 👍点赞✍评论⭐收藏

收录专栏: TensorFlow与Keras实战演绎机器学习

希望政安晨的博客能够对您有所裨益,如有不足之处,欢迎在评论区提出指正!

在TensorFlow中,fit()是一个非常强大和常用的训练函数,它可以批次地训练模型并监测其性能。虽然fit()提供了很多有用的默认行为,但有时您可能想自定义fit()中发生的操作。

一种自定义fit()的方法是使用回调函数。回调函数是在训练过程中的特定时间点被调用的函数,您可以编写自己的回调函数来执行特定的操作。

TensorFlow提供了许多内置的回调函数,如EarlyStoppingCallback、ModelCheckpointCallback等,您也可以编写自己的自定义回调函数。通过将回调函数传递给fit()函数的callbacks参数,您可以自定义在每个训练批次或训练周期结束时发生的操作。

另一种自定义fit()的方法是编写自定义的训练循环。

默认情况下,fit()函数使用单个步骤来执行训练循环,但您可以重写这个步骤来实现自己的训练逻辑。您可以使用TensorFlow的GradientTape来手动计算梯度,并使用优化器来更新模型的权重。通过这种方式,您可以完全控制训练过程中的每个步骤。

最后,您还可以自定义fit()的行为通过设置fit()函数的其他参数。

例如,您可以通过设置batch_size参数来定义每个训练批次的大小,或者通过设置epochs参数来指定训练周期的数量。除了这些参数,您还可以设置其他参数,如学习率、损失函数等,以进一步自定义fit()的行为。

总而言之,TensorFlow提供了多种方法来自定义fit()函数中发生的操作。无论您是通过回调函数、自定义训练循环还是设置fit()的参数,您都可以根据自己的需求来定制训练过程。

这使得TensorFlow成为一个非常灵活和强大的深度学习框架。

今天我们讲的就是keras api 在这里的应用方法


前言

当你进行监督学习时,你可以使用fit(),一切都运行得很顺利。

当你需要控制每一个细节时,你可以完全从头开始编写自己的训练循环。

但是如果你需要一个自定义的训练算法,但仍然希望从fit()的便利功能中受益,比如回调函数,内置的分发支持或步骤融合,该怎么办呢?

Keras的一个核心原则是逐步揭示复杂性。您应该总是能够逐渐进入更低级的工作流程。如果高级功能与您的使用情况不完全匹配,您不应该感到掉入悬崖。您应该能够在保留相应数量高级便利的同时获得对细节的更多控制。

当你需要自定义fit()函数的行为时,你应该重写Model类的训练步骤函数。这是fit()函数在处理每个数据批次时调用的函数。然后,你仍然可以像往常一样调用fit()函数 - 它将会运行你自己的学习算法。

请注意这种模式并不妨碍您使用函数式API构建模型。无论您是构建顺序模型、函数式API模型还是子类化模型,都可以使用这种模式。

现在让我们看看这一切都是怎么工作的吧。

导入

import os# This guide can only be run with the TF backend.
os.environ["KERAS_BACKEND"] = "tensorflow"import tensorflow as tf
import keras
from keras import layers
import numpy as np

来一个简单例子

我们创建一个新的类,继承自keras.Model。

我们只需重写train_step(self, data)方法。

我们返回一个将指标名称(包括损失)映射到它们当前值的字典。

输入参数data是传递给fit作为训练数据的内容:

如果你通过调用fit(x, y, ...)传递NumPy数组,则data将是元组(x, y)

如果你通过调用fit(dataset, ...)传递tf.data.Dataset,则data将是每个批次由数据集产生的内容。

在train_step()方法的主体中,我们实现了一个常规的训练更新,与你已经熟悉的类似。重要的是,我们通过self.compute_loss()计算损失,它包装了传递给compile()的损失函数。

类似地,我们对self.metrics中的指标调用metric.update_state(y, y_pred)来更新在compile()中传递的指标的状态,并在最后从self.metrics查询结果以获取它们的当前值。

class CustomModel(keras.Model):def train_step(self, data):# Unpack the data. Its structure depends on your model and# on what you pass to `fit()`.x, y = datawith tf.GradientTape() as tape:y_pred = self(x, training=True)  # Forward pass# Compute the loss value# (the loss function is configured in `compile()`)loss = self.compute_loss(y=y, y_pred=y_pred)# Compute gradientstrainable_vars = self.trainable_variablesgradients = tape.gradient(loss, trainable_vars)# Update weightsself.optimizer.apply(gradients, trainable_vars)# Update metrics (includes the metric that tracks the loss)for metric in self.metrics:if metric.name == "loss":metric.update_state(loss)else:metric.update_state(y, y_pred)# Return a dict mapping metric names to current valuereturn {m.name: m.result() for m in self.metrics}

现在让咱们试试:

# Construct and compile an instance of CustomModel
inputs = keras.Input(shape=(32,))
outputs = keras.layers.Dense(1)(inputs)
model = CustomModel(inputs, outputs)
model.compile(optimizer="adam", loss="mse", metrics=["mae"])# Just use `fit` as usual
x = np.random.random((1000, 32))
y = np.random.random((1000, 1))
model.fit(x, y, epochs=3)

下沉到更低的级别

当然,你可以在compile()中跳过传递损失函数,而是在train_step中手动完成所有操作。同样,在指标方面也是如此。

这是一个低级示例,只使用compile()配置优化器:

我们首先创建度量实例来跟踪我们的损失和MAE得分(在__init__()中)。

我们实现了一个自定义的train_step()函数,该函数通过调用update_state()来更新这些度量的状态,然后通过调用result()来查询它们的当前平均值,以便在进度条中显示并传递给任何回调函数。

注意,我们需要在每个epoch之间调用reset_states()来重置我们的度量!

否则,调用result()将返回训练开始以来的平均值,而我们通常使用每个epoch的平均值。

幸运的是,框架可以为我们做到这一点:只需将要重置的任何度量列在模型的metrics属性中。

模型将在每次fit()的epoch开始或调用evaluate()的开始时调用此处列出的任何对象的reset_states()函数。

class CustomModel(keras.Model):def __init__(self, *args, **kwargs):super().__init__(*args, **kwargs)self.loss_tracker = keras.metrics.Mean(name="loss")self.mae_metric = keras.metrics.MeanAbsoluteError(name="mae")self.loss_fn = keras.losses.MeanSquaredError()def train_step(self, data):x, y = datawith tf.GradientTape() as tape:y_pred = self(x, training=True)  # Forward pass# Compute our own lossloss = self.loss_fn(y, y_pred)# Compute gradientstrainable_vars = self.trainable_variablesgradients = tape.gradient(loss, trainable_vars)# Update weightsself.optimizer.apply(gradients, trainable_vars)# Compute our own metricsself.loss_tracker.update_state(loss)self.mae_metric.update_state(y, y_pred)return {"loss": self.loss_tracker.result(),"mae": self.mae_metric.result(),}@propertydef metrics(self):# We list our `Metric` objects here so that `reset_states()` can be# called automatically at the start of each epoch# or at the start of `evaluate()`.return [self.loss_tracker, self.mae_metric]# Construct an instance of CustomModel
inputs = keras.Input(shape=(32,))
outputs = keras.layers.Dense(1)(inputs)
model = CustomModel(inputs, outputs)# We don't pass a loss or metrics here.
model.compile(optimizer="adam")# Just use `fit` as usual -- you can use callbacks, etc.
x = np.random.random((1000, 32))
y = np.random.random((1000, 1))
model.fit(x, y, epochs=5)

支持样本权重和类别权重

您可能已经注意到,我们的第一个基本示例没有提及样本加权。

如果您想支持fit()函数的sample_weight和class_weight参数,只需按照以下步骤操作:

从data参数中解包sample_weight 将其传递给compute_loss和update_state函数(当然,如果您不依赖compile()函数来计算损失和指标,您也可以手动应用它) 就是这样。

class CustomModel(keras.Model):def train_step(self, data):# Unpack the data. Its structure depends on your model and# on what you pass to `fit()`.if len(data) == 3:x, y, sample_weight = dataelse:sample_weight = Nonex, y = datawith tf.GradientTape() as tape:y_pred = self(x, training=True)  # Forward pass# Compute the loss value.# The loss function is configured in `compile()`.loss = self.compute_loss(y=y,y_pred=y_pred,sample_weight=sample_weight,)# Compute gradientstrainable_vars = self.trainable_variablesgradients = tape.gradient(loss, trainable_vars)# Update weightsself.optimizer.apply(gradients, trainable_vars)# Update the metrics.# Metrics are configured in `compile()`.for metric in self.metrics:if metric.name == "loss":metric.update_state(loss)else:metric.update_state(y, y_pred, sample_weight=sample_weight)# Return a dict mapping metric names to current value.# Note that it will include the loss (tracked in self.metrics).return {m.name: m.result() for m in self.metrics}# Construct and compile an instance of CustomModel
inputs = keras.Input(shape=(32,))
outputs = keras.layers.Dense(1)(inputs)
model = CustomModel(inputs, outputs)
model.compile(optimizer="adam", loss="mse", metrics=["mae"])# You can now use sample_weight argument
x = np.random.random((1000, 32))
y = np.random.random((1000, 1))
sw = np.random.random((1000, 1))
model.fit(x, y, sample_weight=sw, epochs=3)

提供您自己的评估步骤

如果你想对model.evaluate()的调用做同样的操作,那么你可以通过同样的方式覆盖test_step。下面是实现的样例:

class CustomModel(keras.Model):def test_step(self, data):# Unpack the datax, y = data# Compute predictionsy_pred = self(x, training=False)# Updates the metrics tracking the lossloss = self.compute_loss(y=y, y_pred=y_pred)# Update the metrics.for metric in self.metrics:if metric.name == "loss":metric.update_state(loss)else:metric.update_state(y, y_pred)# Return a dict mapping metric names to current value.# Note that it will include the loss (tracked in self.metrics).return {m.name: m.result() for m in self.metrics}# Construct an instance of CustomModel
inputs = keras.Input(shape=(32,))
outputs = keras.layers.Dense(1)(inputs)
model = CustomModel(inputs, outputs)
model.compile(loss="mse", metrics=["mae"])# Evaluate with our custom test_step
x = np.random.random((1000, 32))
y = np.random.random((1000, 1))
model.evaluate(x, y)

总结:一个端到端的GAN示例

让我们通过一个端到端的示例来演示你刚刚学到的所有内容。

让我们考虑以下内容:

一个生成器网络,用于生成28x28x1的图像。

一个判别器网络,用于将28x28x1的图像分为两类("假"和"真")。

每个网络都有一个优化器。 一个损失函数,用于训练判别器。

# Create the discriminator
discriminator = keras.Sequential([keras.Input(shape=(28, 28, 1)),layers.Conv2D(64, (3, 3), strides=(2, 2), padding="same"),layers.LeakyReLU(negative_slope=0.2),layers.Conv2D(128, (3, 3), strides=(2, 2), padding="same"),layers.LeakyReLU(negative_slope=0.2),layers.GlobalMaxPooling2D(),layers.Dense(1),],name="discriminator",
)# Create the generator
latent_dim = 128
generator = keras.Sequential([keras.Input(shape=(latent_dim,)),# We want to generate 128 coefficients to reshape into a 7x7x128 maplayers.Dense(7 * 7 * 128),layers.LeakyReLU(negative_slope=0.2),layers.Reshape((7, 7, 128)),layers.Conv2DTranspose(128, (4, 4), strides=(2, 2), padding="same"),layers.LeakyReLU(negative_slope=0.2),layers.Conv2DTranspose(128, (4, 4), strides=(2, 2), padding="same"),layers.LeakyReLU(negative_slope=0.2),layers.Conv2D(1, (7, 7), padding="same", activation="sigmoid"),],name="generator",
)

这是一个完整的GAN类,它重写了compile()方法以使用自己的签名,并在train_step中使用17行代码实现了整个GAN算法。

class GAN(keras.Model):def __init__(self, discriminator, generator, latent_dim):super().__init__()self.discriminator = discriminatorself.generator = generatorself.latent_dim = latent_dimself.d_loss_tracker = keras.metrics.Mean(name="d_loss")self.g_loss_tracker = keras.metrics.Mean(name="g_loss")self.seed_generator = keras.random.SeedGenerator(1337)@propertydef metrics(self):return [self.d_loss_tracker, self.g_loss_tracker]def compile(self, d_optimizer, g_optimizer, loss_fn):super().compile()self.d_optimizer = d_optimizerself.g_optimizer = g_optimizerself.loss_fn = loss_fndef train_step(self, real_images):if isinstance(real_images, tuple):real_images = real_images[0]# Sample random points in the latent spacebatch_size = tf.shape(real_images)[0]random_latent_vectors = keras.random.normal(shape=(batch_size, self.latent_dim), seed=self.seed_generator)# Decode them to fake imagesgenerated_images = self.generator(random_latent_vectors)# Combine them with real imagescombined_images = tf.concat([generated_images, real_images], axis=0)# Assemble labels discriminating real from fake imageslabels = tf.concat([tf.ones((batch_size, 1)), tf.zeros((batch_size, 1))], axis=0)# Add random noise to the labels - important trick!labels += 0.05 * keras.random.uniform(tf.shape(labels), seed=self.seed_generator)# Train the discriminatorwith tf.GradientTape() as tape:predictions = self.discriminator(combined_images)d_loss = self.loss_fn(labels, predictions)grads = tape.gradient(d_loss, self.discriminator.trainable_weights)self.d_optimizer.apply(grads, self.discriminator.trainable_weights)# Sample random points in the latent spacerandom_latent_vectors = keras.random.normal(shape=(batch_size, self.latent_dim), seed=self.seed_generator)# Assemble labels that say "all real images"misleading_labels = tf.zeros((batch_size, 1))# Train the generator (note that we should *not* update the weights# of the discriminator)!with tf.GradientTape() as tape:predictions = self.discriminator(self.generator(random_latent_vectors))g_loss = self.loss_fn(misleading_labels, predictions)grads = tape.gradient(g_loss, self.generator.trainable_weights)self.g_optimizer.apply(grads, self.generator.trainable_weights)# Update metrics and return their value.self.d_loss_tracker.update_state(d_loss)self.g_loss_tracker.update_state(g_loss)return {"d_loss": self.d_loss_tracker.result(),"g_loss": self.g_loss_tracker.result(),}

让我们试试吧:

# Prepare the dataset. We use both the training & test MNIST digits.
batch_size = 64
(x_train, _), (x_test, _) = keras.datasets.mnist.load_data()
all_digits = np.concatenate([x_train, x_test])
all_digits = all_digits.astype("float32") / 255.0
all_digits = np.reshape(all_digits, (-1, 28, 28, 1))
dataset = tf.data.Dataset.from_tensor_slices(all_digits)
dataset = dataset.shuffle(buffer_size=1024).batch(batch_size)gan = GAN(discriminator=discriminator, generator=generator, latent_dim=latent_dim)
gan.compile(d_optimizer=keras.optimizers.Adam(learning_rate=0.0003),g_optimizer=keras.optimizers.Adam(learning_rate=0.0003),loss_fn=keras.losses.BinaryCrossentropy(from_logits=True),
)# To limit the execution time, we only train on 100 batches. You can train on
# the entire dataset. You will need about 20 epochs to get nice results.
gan.fit(dataset.take(100), epochs=1)

深度学习背后的思想很简单吧。


本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若转载,请注明出处:http://xiahunao.cn/news/2905041.html

如若内容造成侵权/违法违规/事实不符,请联系瞎胡闹网进行投诉反馈,一经查实,立即删除!

相关文章

【算法题】三道题理解算法思想--滑动窗口篇

滑动窗口 本篇文章中会带大家从零基础到学会利用滑动窗口的思想解决算法题,我从力扣上筛选了三道题,难度由浅到深,会附上题目链接以及算法原理和解题代码,希望大家能坚持看完,绝对能有收获,大家有更好的思…

Jackson 2.x 系列【6】注解大全篇二

有道无术,术尚可求,有术无道,止于术。 本系列Jackson 版本 2.17.0 源码地址:https://gitee.com/pearl-organization/study-jaskson-demo 文章目录 注解大全2.11 JsonValue2.12 JsonKey2.13 JsonAnySetter2.14 JsonAnyGetter2.15 …

【嵌入式机器学习开发实战】(十二)—— 政安晨:通过ARM-Linux掌握基本技能【C语言程序的安装运行】

政安晨的个人主页:政安晨 欢迎 👍点赞✍评论⭐收藏 收录专栏: 嵌入式机器学习开发实战 希望政安晨的博客能够对您有所裨益,如有不足之处,欢迎在评论区提出指正! 在ARM-Linux系统中,C语言程序的安装和运行可…

Yarn简介及Windows安装与使用指南

🌟 前言 欢迎来到我的技术小宇宙!🌌 这里不仅是我记录技术点滴的后花园,也是我分享学习心得和项目经验的乐园。📚 无论你是技术小白还是资深大牛,这里总有一些内容能触动你的好奇心。🔍 &#x…

RTSP应用:实现视频流的实时推送

在实现实时视频流推送的项目中,RTSP(Real Time Streaming Protocol)协议扮演着核心角色。本文将指导你通过安装FFmpeg软件,下载并编译live555,以及配置ffmpeg进行视频流推送,来实现一个基本的RTSP流媒体服务…

Synchronized锁、公平锁、悲观锁乐观锁、死锁等

悲观锁 认为自己在使用数据的时候一定会有别的线程来修改数据,所以在获取数据前会加锁,确保不会有别的线程来修改 如: Synchronized和Lock锁 适合写操作多的场景 乐观锁 适合读操作多的场景 总结: 线程8锁🔐 调用 声明 结果:先打印发送短信,后打印发送邮件 结论…

网络:udptcp套接字

目录 协议 网络传输基本流程 网络编程套接字 udp套接字编程 udp相关代码实现 sock函数 bind函数 recvfrom函数 sendto函数 udp执行指令代码 popen函数 udp多线程版收发消息 tcp套接字编程 tcp套接字代码 listen函数 accept函数 read/write函数 connect函数 recv/…

计算机网络——29ISP之间的路由选择:BGP

ISP之间的路由选择:BGP 层次路由 一个平面的路由 一个网络中的所有路由器的地位一样通过LS,DV,或者其他路由算法,所有路由器都要知道其他所有路由器(子网)如何走所有路由器在一个平面 平面路由的问题 …

数据结构与算法 双链表的转置

一、实验内容 有一个带头结点的双链表L,设计一个算法将其所有元素逆置,即第一个元素变为最后一个元素,第2个元素变为倒数第2个元素,最后一个元素变为第1个元素。 二、实验步骤 1、dlinklist.cpp 2、reverse.cpp 三、实验结果 四…

JAVA 源码分析Integer的128陷阱

128陷阱介绍及演示 首先什么是128陷阱? Integer包装类两个值大小在-128到127之间时可以判断两个数相等,因为两个会公用同一个对象,返回true, 但是超过这个范围两个数就会不等,因为会变成两个对象,返回fal…

《Vision mamba》论文笔记

原文出处: [2401.09417] Vision Mamba: Efficient Visual Representation Learning with Bidirectional State Space Model (arxiv.org) 原文笔记: What: Vision Mamba: Efficient Visual Representation Learning with Bidirectional St…

啥也不会的大学生看过来,这8步就能系统入门stm32单片机???

大家好,今天给大家介绍啥也不会的大学生看过来,这8步就能系统入门stm32单片机,文章末尾附有分享大家一个资料包,差不多150多G。里面学习内容、面经、项目都比较新也比较全!可进群免费领取。 对于没有任何基础的大学生来…

HTTP状态 405 - 方法不允许

方法有问题。 用Post发的请求&#xff0c;然后用Put接收的。 大家也可以看看是不是有这种问题 <body><h1>HTTP状态 405 - 方法不允许</h1><hr class"line" /><p><b>类型</b> 状态报告</p><p><b>消息…

状态模式实战运用

目录 前言 UML plantuml 类图 实战代码 Form State Client 前言 通常一个完整的业务流程中&#xff0c;会经历多个阶段&#xff0c;每个阶段即一个业务状态&#xff0c;不同状态下对应这不同的业务处理逻辑。 无脑堆砌 if else 做判断然后选择对应的业务处理其实也能…

【MySQL】6.MySQL主从复制和读写分离

主从复制 主从复制与读写分离 通常数据库的读/写都在同一个数据库服务器中进行&#xff1b; 但这样在安全性、高可用性和高并发等各个方面无法满足生产环境的实际需求&#xff1b; 因此&#xff0c;通过主从复制的方式同步数据&#xff0c;再通过读写分离提升数据库的并发负载…

Adobe最近推出了Firefly AI的结构参考以及面向品牌的GenStudio

每周跟踪AI热点新闻动向和震撼发展 想要探索生成式人工智能的前沿进展吗&#xff1f;订阅我们的简报&#xff0c;深入解析最新的技术突破、实际应用案例和未来的趋势。与全球数同行一同&#xff0c;从行业内部的深度分析和实用指南中受益。不要错过这个机会&#xff0c;成为AI领…

数据结构七大常见的排序

数据结构七大常见的排序 常见排序算法分类1.插入排序2.希尔排序(缩小增量排序)3.选择排序4.堆排序5.冒泡排序6.快速排序7.归并排序 常见排序算法分类 1.插入排序 基本思想&#xff1a;把待排序的数组按大小逐个插入到一个已经排好序的有序序列中&#xff0c;直到所有的数据插入…

Django 评论楼创建

Django 评论楼创建 【零】最终效果预览 【一】介绍 &#xff08;1&#xff09;情况说明 在Django模型层中有这么个字段 parent models.ForeignKey(toself, on_deletemodels.CASCADE, verbose_name"父评论ID", nullTrue, blankTrue)这个字段是一对多的外键字段 其…

linux中查看内存占用空间

文章目录 linux中查看内存占用空间 linux中查看内存占用空间 使用 df -h 查看磁盘空间 使用 du -sh * 查看每个目录的大小 注意这里是当前目录下的文件大小&#xff0c;查看系统的可以回到根目录 经过查看没有发现任何大的文件夹。 继续下面的步骤 如果您的Linux磁盘已满&a…

快速上手Spring Cloud 十五:与人工智能的智慧交融

快速上手Spring Cloud 一&#xff1a;Spring Cloud 简介 快速上手Spring Cloud 二&#xff1a;核心组件解析 快速上手Spring Cloud 三&#xff1a;API网关深入探索与实战应用 快速上手Spring Cloud 四&#xff1a;微服务治理与安全 快速上手Spring Cloud 五&#xff1a;Spring …