深度学习——keras中的Sequential和Functional API

2024-03-27 09:38

本文主要是介绍深度学习——keras中的Sequential和Functional API,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

大家好,上期推文介绍了Keras的一些特点和一些基本的知识点,不知道大家在平时的时间有没有自己学习一下深度学习相关的知识,或者机器学习相关的知识呢?有这些的预备知识对于学这个专题还是有帮助的。

本期内容我们先来聊一聊Keras中模型的种类,也就是来聊聊Sequential模型与Functional模型,即序贯模型和函数式模型,我们一个个来看。

 

一、Sequential模型

神经网络模型是一种将信息朝着某一个方向进行传递的模型,方向性的传递形式就很适合以一种顺序(序贯)的数据结构来进行表示,有一种“一条路走到黑”的感觉。在工程项目中使用序贯模型可以解决很多的实际需求,它的主要特点如下:

  1. 简单的线性结构

  2. 从开始到结束的结构顺序

  3. 没有分叉结构

  4. 多个网络层(输入层,隐层,激活层等等)的线性堆叠

不要觉得它简单功能就不强大,要知道全连接神经网络、卷积神经网络(CNN)、循环神经网络(RNN)等网络都可以使用Sequential Model来进行构建。其构建

模型通常分为五个步骤: 1.定义模型 2.定义优化目标 3.输入数据 4.训练模型 5.评估模型的性能

这里要提醒的是,大家在学习Keras中的时候要有层(layers)的概念。如卷积层,池化层,全连接层,LSTM层等等。在Keras中我们构建模型可以直接使用这些层来构建一个目标神经网络。接下来我们简单的来说说这几个步骤:

 

1.1 定义模型

可以通过向Sequential模型传递一个层的列表来构造该模型:

from keras.models import Sequential
from keras.layers import Dense, Activation
model = Sequential([Dense(32,  input_shape=(784,)),Activation('relu'),Dense(10),Activation('softmax'),])

上述代码是使用层的列表来构建Sequential模型的,其中堆叠了四层:

第一层是全连接层,神经元个数为32,input_shape=(784,);

第二层是激活层,激活函数为relu;

第三层是全连接层,神经元个数为10;

第二层是激活层,激活函数为softmax;

当然了,我们也可以使用add()方法进行层的堆叠:

model = Sequential()
model.add(Dense(32, input_shape=(784,)))
model.add(Activation('relu'))
model.add(Dense(10))
model.add(Activation('softmax'))

个人比较喜欢第二种方式,十分简单

大家可能有一个疑问,在第一层使用了input_shape,其他层却没有传入这个参数,这是为什么呢?这是因为:后面的各个层可以自动的推导出中间数据的shape,因此就不用传参了。

 

1.2 定义优化目标

我们知道不同种类的任务有不同的优化目标,也就存在要使用不同的优化器(Optimizer)、损失函数(Loss)和评估标准(Metrics)。在实际的项目工程中,多分类的损失函数我们一般选择categorical_crossentropy,回归问题常用MSE损失函数。

其实这个定义优化目标的过程是整个构建模型中最重要的部分,因为在这部分我们要选择优化器和损失函数。通过选择的优化器和损失函数来对模型中的参数进行计算和更新。在Keras中我们是使用compile()方法来完成:

# 多分类问题
model.compile(optimizer='rmsprop',loss='categorical_crossentropy',metrics=['accuracy'])
# 回归问题
model.compile(optimizer='rmsprop',loss='mse')

1.3 输入数据和训练模型

定义优化目标以后就需要进行数据的输入了,在Keras中是使用fit()方法将数据传递到构建的模型当中的。fit()方法具有较多的参数,其返回值是一个History类对象,这个对象包含两个属性,分别为epoch和history,epoch为训练轮数,history字典类型,包含val_loss,val_acc,loss,acc四个key值。

大家可以查看到这些参数的具体含义。一般情况下我们使用的参数主要有:x, y, batch_size,epochs,verbose,validation_split, validation_data等。如:

model.fit(X_train, y_train, epochs=10, batch_size=64, verbose=1, validation_split=0.05)

上述代码表示我们输入了X_train, y_train训练集数据,迭代次数为10,数据批次大小为64,verbose = 1表示输出进度条记录,validation_split=0.05表示从训练集中取出0.05比例的数据作为验证集。

 

1.4 评估模型性能

fit()之后就是在测试数据集上对我们的数据进行验证了,在Keras中我们可以这样来进行模型的评估:

score = model.evaluate(X_test, y_test, batch_size = 64)

这个模型的评估方法使用很简单。大家可以看上篇推文,推文中的手写字体识别模型就是使用Sequential模型来搭建的,代码我也贴一下:

from keras.datasets import mnist
from keras.utils import np_utils
from keras.models import Sequential
from keras.layers import Dense,Activation
path = r'C:\Users\LEGION\Desktop\datasets\mnist.npz'
(X_train, y_train), (X_test, y_test) = mnist.load_data(path)X_train = X_train.reshape(len(X_train),-1)
X_test = X_test.reshape(len(X_test), -1)
X_train = X_train.astype('float32')/255
X_test = X_test.astype('float32')/255y_train = np_utils.to_categorical(y_train)
y_test = np_utils.to_categorical(y_test)model = Sequential()
model.add(Dense(512, input_shape=(28*28,),activation='relu'))
model.add(Dense(10,activation='softmax'))
model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])
model.fit(X_train, y_train, epochs=10, batch_size=64, verbose=1, validation_split=0.05)
loss, accuracy = model.evaluate(X_test, y_test)
Testloss, Testaccuracy = model.evaluate(X_test, y_test)
print('Testloss:', Testloss)
print('Testaccuracy:', Testaccuracy)

 

二、Functional模型

接下来我们来看一下Functional模型。函数式模型称作Functional,仔细研究发现其类名是Model,因此有时候也用Model来代表函数式模型,即Model模型。

我们知道Sequential模型类似一条路走到㡳的结构,如VGG也是一条路走到㡳的模型。那么在实际的工程项目中,如果我们的目标是多个输出或者非循环有向模型,此时我们就应该选择使用函数式模型来构建。

函数式模型是最广泛的一类模型,Sequential模型是它的特例,下图给出了几种模型的简单示例:

Sequential模型示例

多输入单输出Functional模型的示例

单输入多输出Functional模型的示例

Keras 的函数式模型把层(layers)当作函数来使用,接收张量并返回张量。

我们先使用函数式模型来实现一个序贯模型(全连接网络模型):

from keras import Input
from keras import layers
from keras import models# 构建输入层的张量
input_tensor = Input(shape=(16,))
# 将输入层的返回的input_tensor作为参数输入到第一层网络的Dense中
dense_1 = layers.Dense(32, activation='relu')(input_tensor)
# 将第一层的返回的dense_1作为参数输入到第二层的Dense中
dense_2 = layers.Dense(64, activation='relu')(dense_1)
# 最后一层来构建输出层,并使用10个神经元来进行分类
output_tensor = layers.Dense(10, activation='softmax')(dense_2)
# 将输入层和输出层的张量输入模型中
model = models.Model(input_tensor, output_tensor)# 显示模型的相关信息
model.summary()

上述模型的网络结构示意图如下:

我们从图中可以了解到每一层的输出形状,参数情况。

那么基于Functional模型构建过程一般的步骤如下:

  1. 首先使用Keras的Input()方法来构建输入层,该层一般不含可输入的参数。

  2. 将输入层返回的张量继续输入到下一层,如果还有下一层,继续此操作

  3. 将输入层和输出层的张量输入模型中

 

上述代码只是使用函数式模型来构建一个非常简单的全连接网络。我们还未实现多输入单输出,单输入多输出和多输入多输出的模型,后期打算一个种类使用一期推文来进行详细的介绍。大家可以先学一下机器学习中的多模型融合相关的技术,这对于后期的学习十分有利。

 

三、总结

想必大家也已经知道了一些Sequential顺序模型和Functional模型相关的知识。对于Sequential而言掌握它是很容易的,大家或许对于函数式模型还不是很熟悉,但是没关系,后续我们会详细的探究不同输入和输出形式的函数式模型,到时候也会结合实例来进行阐述。

这篇关于深度学习——keras中的Sequential和Functional API的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



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

相关文章

PHP应用中处理限流和API节流的最佳实践

《PHP应用中处理限流和API节流的最佳实践》限流和API节流对于确保Web应用程序的可靠性、安全性和可扩展性至关重要,本文将详细介绍PHP应用中处理限流和API节流的最佳实践,下面就来和小编一起学习... 目录限流的重要性在 php 中实施限流的最佳实践使用集中式存储进行状态管理(如 Redis)采用滑动

深度解析Python中递归下降解析器的原理与实现

《深度解析Python中递归下降解析器的原理与实现》在编译器设计、配置文件处理和数据转换领域,递归下降解析器是最常用且最直观的解析技术,本文将详细介绍递归下降解析器的原理与实现,感兴趣的小伙伴可以跟随... 目录引言:解析器的核心价值一、递归下降解析器基础1.1 核心概念解析1.2 基本架构二、简单算术表达

深度解析Java @Serial 注解及常见错误案例

《深度解析Java@Serial注解及常见错误案例》Java14引入@Serial注解,用于编译时校验序列化成员,替代传统方式解决运行时错误,适用于Serializable类的方法/字段,需注意签... 目录Java @Serial 注解深度解析1. 注解本质2. 核心作用(1) 主要用途(2) 适用位置3

Java MCP 的鉴权深度解析

《JavaMCP的鉴权深度解析》文章介绍JavaMCP鉴权的实现方式,指出客户端可通过queryString、header或env传递鉴权信息,服务器端支持工具单独鉴权、过滤器集中鉴权及启动时鉴权... 目录一、MCP Client 侧(负责传递,比较简单)(1)常见的 mcpServers json 配置

Maven中生命周期深度解析与实战指南

《Maven中生命周期深度解析与实战指南》这篇文章主要为大家详细介绍了Maven生命周期实战指南,包含核心概念、阶段详解、SpringBoot特化场景及企业级实践建议,希望对大家有一定的帮助... 目录一、Maven 生命周期哲学二、default生命周期核心阶段详解(高频使用)三、clean生命周期核心阶

深度剖析SpringBoot日志性能提升的原因与解决

《深度剖析SpringBoot日志性能提升的原因与解决》日志记录本该是辅助工具,却为何成了性能瓶颈,SpringBoot如何用代码彻底破解日志导致的高延迟问题,感兴趣的小伙伴可以跟随小编一起学习一下... 目录前言第一章:日志性能陷阱的底层原理1.1 日志级别的“双刃剑”效应1.2 同步日志的“吞吐量杀手”

Unity新手入门学习殿堂级知识详细讲解(图文)

《Unity新手入门学习殿堂级知识详细讲解(图文)》Unity是一款跨平台游戏引擎,支持2D/3D及VR/AR开发,核心功能模块包括图形、音频、物理等,通过可视化编辑器与脚本扩展实现开发,项目结构含A... 目录入门概述什么是 UnityUnity引擎基础认知编辑器核心操作Unity 编辑器项目模式分类工程

深度解析Python yfinance的核心功能和高级用法

《深度解析Pythonyfinance的核心功能和高级用法》yfinance是一个功能强大且易于使用的Python库,用于从YahooFinance获取金融数据,本教程将深入探讨yfinance的核... 目录yfinance 深度解析教程 (python)1. 简介与安装1.1 什么是 yfinance?

Python学习笔记之getattr和hasattr用法示例详解

《Python学习笔记之getattr和hasattr用法示例详解》在Python中,hasattr()、getattr()和setattr()是一组内置函数,用于对对象的属性进行操作和查询,这篇文章... 目录1.getattr用法详解1.1 基本作用1.2 示例1.3 原理2.hasattr用法详解2.

Go语言使用net/http构建一个RESTful API的示例代码

《Go语言使用net/http构建一个RESTfulAPI的示例代码》Go的标准库net/http提供了构建Web服务所需的强大功能,虽然众多第三方框架(如Gin、Echo)已经封装了很多功能,但... 目录引言一、什么是 RESTful API?二、实战目标:用户信息管理 API三、代码实现1. 用户数据