解码 ResNet:残差块如何增强深度学习性能【数学推导】

2024-06-18 16:28

本文主要是介绍解码 ResNet:残差块如何增强深度学习性能【数学推导】,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

ResNet简介

残差网络结构

残差网络(ResNet)是由何凯明等人在2015年提出的,它极大地提高了深度神经网络的训练效果,尤其是非常深的网络。ResNet的核心思想是引入“残差块”(Residual Block),通过跳跃连接(Shortcut Connection)解决深层网络的梯度消失和梯度爆炸问题。

结构示意图

  • 输入层
  • 一系列的卷积层(Conv Layers)
  • 残差块(Residual Blocks)
  • 全连接层(Fully Connected Layer)
  • 输出层

在传统的卷积神经网络中,每一层都会对输入的特征进行某种变换,比如卷积操作,然后直接输出这些变换后的结果到下一层。可以把这种变换看作是对输入进行处理和提取新的特征。
y l = F l ( x l ) \mathbf{y}_l = \mathcal{F}_l(\mathbf{x}_l) yl=Fl(xl)

而ResNet通过增加一条跳跃连接,使得每个残差块输出的是“变换后的特征+原始输入特征”,即:

y = F ( x , { W i } ) + x \mathbf{y} = \mathcal{F}(\mathbf{x}, \{W_i\}) + \mathbf{x} y=F(x,{Wi})+x

其中, F ( x , { W i } ) \mathcal{F}(\mathbf{x}, \{W_i\}) F(x,{Wi}) 表示通过多层卷积、激活等操作后的特征, x \mathbf{x} x 表示原始输入特征。

什么是跳跃连接?

跳跃连接(Shortcut Connection),又称为“短路连接”或“直连”,是一种直接将输入信号传递到输出信号的技术。具体来说,就是在每个残差块中,除了正常的变换路径外,还增加了一条直接连接输入和输出的路径。

为什么要使用跳跃连接?

在深层网络中,随着层数的增加,梯度可能会逐渐消失或者爆炸,这会导致网络很难训练。而跳跃连接的引入可以缓解这个问题,因为它允许梯度直接传递到前面的层,确保梯度不会消失。

跳跃连接如何缓解梯度消失和梯度爆炸问题

为了理解跳跃连接如何缓解梯度消失和梯度爆炸问题,我们需要从反向传播(Backpropagation)的角度分析梯度传递过程。

在传统的深层网络中,假设某一层的输入是 x l \mathbf{x}_l xl ,输出是 y l \mathbf{y}_l yl 。每层的变换函数记为 F l \mathcal{F}_l Fl,那么:

y l = F l ( x l ) \mathbf{y}_l = \mathcal{F}_l(\mathbf{x}_l) yl=Fl(xl)

而在ResNet中,增加了跳跃连接后,输出变为:

y l = F l ( x l ) + x l \mathbf{y}_l = \mathcal{F}_l(\mathbf{x}_l) + \mathbf{x}_l yl=Fl(xl)+xl

在反向传播中,我们需要计算每层的梯度。对于传统的深层网络,第 l l l 层的梯度计算如下:

∂ L ∂ x l = ∂ L ∂ y l ⋅ ∂ y l ∂ x l = ∂ L ∂ y l ⋅ ∂ F l ( x l ) ∂ x l \frac{\partial \mathcal{L}}{\partial \mathbf{x}_l} = \frac{\partial \mathcal{L}}{\partial \mathbf{y}_l} \cdot \frac{\partial \mathbf{y}_l}{\partial \mathbf{x}_l} = \frac{\partial \mathcal{L}}{\partial \mathbf{y}_l} \cdot \frac{\partial \mathcal{F}_l(\mathbf{x}_l)}{\partial \mathbf{x}_l} xlL=ylLxlyl=ylLxlFl(xl)

而在ResNet中,由于增加了跳跃连接,梯度的计算变为:

∂ L ∂ x l = ∂ L ∂ y l ⋅ ( ∂ F l ( x l ) ∂ x l + I ) \frac{\partial \mathcal{L}}{\partial \mathbf{x}_l} = \frac{\partial \mathcal{L}}{\partial \mathbf{y}_l} \cdot \left( \frac{\partial \mathcal{F}_l(\mathbf{x}_l)}{\partial \mathbf{x}_l} + \mathbf{I} \right) xlL=ylL(xlFl(xl)+I)

这里, I \mathbf{I} I 是单位矩阵,表示跳跃连接的梯度。

梯度分析

在ResNet中,由于跳跃连接的存在,梯度不仅传递了变换部分( ∂ F l ( x l ) ∂ x l \frac{\partial \mathcal{F}_l(\mathbf{x}_l)}{\partial \mathbf{x}_l} xlFl(xl) ),还传递了输入部分( I \mathbf{I} I ),这意味着即使在深层网络中,梯度也能有效地通过跳跃连接传递到前面的层,而不会完全依赖于 ∂ F l ( x l ) ∂ x l \frac{\partial \mathcal{F}_l(\mathbf{x}_l)}{\partial \mathbf{x}_l} xlFl(xl)

具体来说,如果 ∂ F l ( x l ) ∂ x l \frac{\partial \mathcal{F}_l(\mathbf{x}_l)}{\partial \mathbf{x}_l} xlFl(xl) 在深层网络中趋近于0(梯度消失)或趋近于无穷大(梯度爆炸),跳跃连接的单位矩阵 I \mathbf{I} I 确保了梯度至少能通过 I \mathbf{I} I 进行传递,缓解了梯度消失或爆炸的问题。

总结

  1. 跳跃连接的引入:在每个残差块中,除了对输入特征进行卷积、归一化和激活等操作外,还增加了一条直接传递输入特征到输出的路径。
  2. 公式中的体现:输出特征不仅包含变换后的特征,还加上了输入特征,即 y = F ( x ) + x \mathbf{y} = \mathcal{F}(\mathbf{x}) + \mathbf{x} y=F(x)+x
  3. 缓解梯度问题:跳跃连接确保了梯度在反向传播过程中,即使变换部分的梯度消失或爆炸,输入特征的梯度(\mathbf{I})也能直接传递,避免梯度完全消失或爆炸。

残差块的组成及功能

残差块是ResNet的基本单元,每个残差块中包含了两个主要部分:

  1. 变换路径:对输入进行卷积、批量归一化和激活操作。
  2. 跳跃连接(Shortcut Connection):直接将输入传递到输出,不进行任何变换,只是将输入特征原样添加到经过变换后的特征上。

详细组成

  1. 卷积层(Convolutional Layer):提取特征。
  2. 批量归一化层(Batch Normalization Layer):加速训练,稳定输入。
  3. ReLU激活函数(ReLU Activation Function):引入非线性,提高网络表达能力。
  4. 跳跃连接(Shortcut Connection):将输入直接加到输出上。

具体的操作流程如下:

  1. 输入特征 x \mathbf{x} x 通过卷积层和批量归一化层,得到变换后的特征 F ( x ) \mathcal{F}(\mathbf{x}) F(x)
  2. 变换后的特征 F ( x ) \mathcal{F}(\mathbf{x}) F(x) 与输入特征 x \mathbf{x} x 相加,得到输出特征 y \mathbf{y} y

y = F ( x , { W i } ) + x \mathbf{y} = \mathcal{F}(\mathbf{x}, \{W_i\}) + \mathbf{x} y=F(x,{Wi})+x

这里, x \mathbf{x} x 直接通过跳跃连接加到变换后的特征 F ( x ) \mathcal{F}(\mathbf{x}) F(x) 上。

  1. 输出特征 y \mathbf{y} y 再经过ReLU激活函数:

y = ReLU ( y ) \mathbf{y} = \text{ReLU}(\mathbf{y}) y=ReLU(y)

这种设计可以确保即使在深层网络中,梯度也能有效传播,避免梯度消失或爆炸。

ResNet的输出计算

在ResNet中,每一层的输出不仅仅取决于当前层的输入,还包括了前面层的输入,这种设计使得网络能够更有效地学习。

详细的数学推导
假设一个简单的ResNet包含L个残差块,每个残差块输出为 y l \mathbf{y}_l yl ,输入为 x l \mathbf{x}_l xl ,则有:

y l = F l ( x l ) + x l \mathbf{y}_l = \mathcal{F}_l(\mathbf{x}_l) + \mathbf{x}_l yl=Fl(xl)+xl

其中 F l ( x l ) \mathcal{F}_l(\mathbf{x}_l) Fl(xl) 表示第l个残差块中的变换函数(例如两层卷积和ReLU激活函数)。

整个网络的输入为 x 0 \mathbf{x}_0 x0 ,输出为 y L \mathbf{y}_L yL,即:

y L = F L ( y L − 1 ) + y L − 1 \mathbf{y}_L = \mathcal{F}_L(\mathbf{y}_{L-1}) + \mathbf{y}_{L-1} yL=FL(yL1)+yL1
y L − 1 = F L − 1 ( y L − 2 ) + y L − 2 \mathbf{y}_{L-1} = \mathcal{F}_{L-1}(\mathbf{y}_{L-2}) + \mathbf{y}_{L-2} yL1=FL1(yL2)+yL2
⋮ \vdots
y 1 = F 1 ( x 0 ) + x 0 \mathbf{y}_1 = \mathcal{F}_1(\mathbf{x}_0) + \mathbf{x}_0 y1=F1(x0)+x0

逐层递推,我们可以得到最终的输出:

y L = x 0 + ∑ l = 1 L F l ( x l ) \mathbf{y}_L = \mathbf{x}_0 + \sum_{l=1}^{L} \mathcal{F}_l(\mathbf{x}_l) yL=x0+l=1LFl(xl)

这种设计可以看作是对输入的逐层增强,每层不仅仅是对输入的简单变换,更是对前面所有层次特征的累积。

总结

  1. 残差网络结构:ResNet引入了残差块,每个残差块中有一条跳跃连接直接将输入加到输出上,这样即使网络很深,信息也能有效传递。
  2. 残差块的组成及功能:每个残差块由卷积、批量归一化、ReLU激活和跳跃连接组成,确保输入信息能够直接加到输出上。
  3. ResNet的输出计算:通过逐层递推,每一层的输出都是对输入和变换后特征的累积,使得网络能够更有效地学习深层特征。

具体实现:残差块的工作原理

  1. 输入特征(原始输入特征):假设输入特征是 x \mathbf{x} x
  2. 变换路径:输入特征 x \mathbf{x} x 经过一系列的卷积操作、批量归一化和激活函数后,得到变换后的特征 F ( x ) \mathcal{F}(\mathbf{x}) F(x)
  3. 跳跃连接:在变换路径之外,直接将输入特征 x \mathbf{x} x 加到变换后的特征 F ( x ) \mathcal{F}(\mathbf{x}) F(x)上,得到输出特征 y \mathbf{y} y

y = F ( x ) + x \mathbf{y} = \mathcal{F}(\mathbf{x}) + \mathbf{x} y=F(x)+x

这里, F ( x ) \mathcal{F}(\mathbf{x}) F(x) 是通过卷积和激活操作后的特征, x \mathbf{x} x 是原始输入特征。这样,每个残差块的输出就是“变换后的特征+原始输入特征”。

这篇关于解码 ResNet:残差块如何增强深度学习性能【数学推导】的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



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

相关文章

Springboot中分析SQL性能的两种方式详解

《Springboot中分析SQL性能的两种方式详解》文章介绍了SQL性能分析的两种方式:MyBatis-Plus性能分析插件和p6spy框架,MyBatis-Plus插件配置简单,适用于开发和测试环... 目录SQL性能分析的两种方式:功能介绍实现方式:实现步骤:SQL性能分析的两种方式:功能介绍记录

Java深度学习库DJL实现Python的NumPy方式

《Java深度学习库DJL实现Python的NumPy方式》本文介绍了DJL库的背景和基本功能,包括NDArray的创建、数学运算、数据获取和设置等,同时,还展示了如何使用NDArray进行数据预处理... 目录1 NDArray 的背景介绍1.1 架构2 JavaDJL使用2.1 安装DJL2.2 基本操

最长公共子序列问题的深度分析与Java实现方式

《最长公共子序列问题的深度分析与Java实现方式》本文详细介绍了最长公共子序列(LCS)问题,包括其概念、暴力解法、动态规划解法,并提供了Java代码实现,暴力解法虽然简单,但在大数据处理中效率较低,... 目录最长公共子序列问题概述问题理解与示例分析暴力解法思路与示例代码动态规划解法DP 表的构建与意义动

Tomcat高效部署与性能优化方式

《Tomcat高效部署与性能优化方式》本文介绍了如何高效部署Tomcat并进行性能优化,以确保Web应用的稳定运行和高效响应,高效部署包括环境准备、安装Tomcat、配置Tomcat、部署应用和启动T... 目录Tomcat高效部署与性能优化一、引言二、Tomcat高效部署三、Tomcat性能优化总结Tom

Go中sync.Once源码的深度讲解

《Go中sync.Once源码的深度讲解》sync.Once是Go语言标准库中的一个同步原语,用于确保某个操作只执行一次,本文将从源码出发为大家详细介绍一下sync.Once的具体使用,x希望对大家有... 目录概念简单示例源码解读总结概念sync.Once是Go语言标准库中的一个同步原语,用于确保某个操

C#使用yield关键字实现提升迭代性能与效率

《C#使用yield关键字实现提升迭代性能与效率》yield关键字在C#中简化了数据迭代的方式,实现了按需生成数据,自动维护迭代状态,本文主要来聊聊如何使用yield关键字实现提升迭代性能与效率,感兴... 目录前言传统迭代和yield迭代方式对比yield延迟加载按需获取数据yield break显式示迭

使用C#代码计算数学表达式实例

《使用C#代码计算数学表达式实例》这段文字主要讲述了如何使用C#语言来计算数学表达式,该程序通过使用Dictionary保存变量,定义了运算符优先级,并实现了EvaluateExpression方法来... 目录C#代码计算数学表达式该方法很长,因此我将分段描述下面的代码片段显示了下一步以下代码显示该方法如

五大特性引领创新! 深度操作系统 deepin 25 Preview预览版发布

《五大特性引领创新!深度操作系统deepin25Preview预览版发布》今日,深度操作系统正式推出deepin25Preview版本,该版本集成了五大核心特性:磐石系统、全新DDE、Tr... 深度操作系统今日发布了 deepin 25 Preview,新版本囊括五大特性:磐石系统、全新 DDE、Tree

Java实现任务管理器性能网络监控数据的方法详解

《Java实现任务管理器性能网络监控数据的方法详解》在现代操作系统中,任务管理器是一个非常重要的工具,用于监控和管理计算机的运行状态,包括CPU使用率、内存占用等,对于开发者和系统管理员来说,了解这些... 目录引言一、背景知识二、准备工作1. Maven依赖2. Gradle依赖三、代码实现四、代码详解五

Node.js 中 http 模块的深度剖析与实战应用小结

《Node.js中http模块的深度剖析与实战应用小结》本文详细介绍了Node.js中的http模块,从创建HTTP服务器、处理请求与响应,到获取请求参数,每个环节都通过代码示例进行解析,旨在帮... 目录Node.js 中 http 模块的深度剖析与实战应用一、引言二、创建 HTTP 服务器:基石搭建(一