突破次元壁 一文解构神经常微分方程
全文共5838字,预计学习时长20分钟或更长
图片来源:unsplash.com/@alinnnaaaa

神经常微分方程学习动力系统的影象
今天,本文将带你回顾2018年度神经资讯处理系统大会(NIPS)中的最佳论文奖:《神经常微分方程》(Neural ODEs)Neural Ordinary Differential Equations。
论文传送门:https://arxiv.org/abs/1806.07366
这篇文章的重点将介绍神经常微分方程的实际用途、应用这种所需的神经网络型别的方式、原因以及可行性。
GitHub程式码传送门:https://github.com/Rachnog/Neural-ODE-Experiments

为什么需要关注常微分方程?
首先,快速回顾一下什么是常微分方程。它描述了某个变数(这就是为什么是常微分)在某个过程中的变化,这种随时间的变化用导数来表示为:
简单的常微分方程例子
如果存在一些初始条件(变化过程的起始点),并且想要观察该过程将如何发展到某个最终状态的话,我们可以探讨此微分方程的求解。函式解也称为积分曲线(因为可以对方程进行积分得到解x(t))。让我们尝试使用SymPy包来求解上图中方程:
from sympy import dsolve, Eq, symbols, Function
t = symbols('t')
x = symbols('x', cls=Function)
deqn1 = Eq(x(t).diff(t), 1 - x(t))
sol1 = dsolve(deqn1, x(t))
则会得出
Eq(x(t), C1*exp(-t) + 1)
其中C1是常数,可以在给定一些初始条件的情况下确定。如果以适当的形式给出,则可以用解析法求解,但通常用数值法求解。最古老、最简单的算法之一是尤拉法。其核心思想是用切线逐步逼近函式解:
http://tutorial.math.lamar.edu/Classes/DE/EulersMe
请访问图片下面的连结以获得更详细的说明。但最后,对于此方程,我们可得出一个非常简单的公式:
http://tutorial.math.lamar.edu/Classes/DE/EulersMe
在n个时间步长的离散网格上的解为:
http://tutorial.math.lamar.edu/Classes/DE/EulersMe

ResNets是常微分方程的解吗?
当然!y_{n+1} = y_n + f(t_n, y_n) 只是ResNet中的一个残差连线,其中某个层的输出是 f()层本身的输出与该层y_n的输入的总和。这基本上是神经常微分方程的主要思想:神经网络中的残差块链基本上是用尤拉法来求解常微分方程。
在这种情况下,系统的初始条件是“时间”0,它表示神经网络的第一层,因为x(0)将作用于正常输入,这可以是时间序列、影象,以及任何你想要的!“时间” t的最终条件是神经网络的期望输出:标量值、表示类的或其他任何东西的向量。
如果这些残差连线是尤拉法的离散时间步长,那么就意味着只要选择离散方案,可以调节神经网络的深度。因此,可以使解(又名神经网络)更精确或更粗略,甚至使它延伸至无限层!
具有固定层数的 ResNet与具有灵活层数的常微分方程网之间的差异
尤拉法是不是过于原始了?确实如此,所以让我们用一些抽象的概念来代替ResNet 或是 EulerSolverNet,比如ODESolveNet,其中ODESolve将是一个函式,它提供了一个比尤拉法更精确的常微分方程(简单地说:神经网络本身)解决方案。网络体系结构现在可能如下所示:
nn = Network(
Dense(...), # making some primary embedding
ODESolve(...), # "infinite-layer neural network"
Dense(...) # output layer
)
我们忘记了一件事……神经网络是一个可微函式,所以我们可以用基于梯度的优化过程来完善它。我们应该如何通过ODESolve() 函式进行反向传播?在例子中,这实际上也像一个黑匣子。特别地,我们需要一个由输入和动力学引数组成的损失函式斜率。这种数学方法叫做伴随敏度法。可参考原始档案(https://arxiv.org/pdf/1806.07366.pdf)和教程(https://nbviewer.jupyter.org/github/urtrial/neural_ode/blob/master/Neural%20ODEs%20%28Russian%29.ipynb)以获得更多详细资讯,但其本质展示于下图中(L代表我们要优化的主要损失函式):
为ODESolve()法设计“反向传播”梯度
简而言之,伴随系统与描述该过程的原始动力系统一起,通过链式规则(即众所周知的反向传播的根源所在)描述后续过程的每个点的导数状态。正是由此可以得到导数的初始状态,并以类似的方式,通过一个函式的引数即动力学建模(一个“残差块”,或“旧”尤拉法的离散步骤)。
建议观看论文作者之一的简报,来获取更多资讯:

神经常微分方程的可能应用
首先,与“正常ResNet”相比的优势和动机:
· 高效率内存:在反向传播时,不需要储存所有引数和变化率。
· 自拟合计算:可用离散化方案来平衡速度和精度,而且,在训练和推理时会有不同的离散化方案。
· 高效率引数:附近“层”的引数自动系结在一起(参考:https://arxiv.org/pdf/1806.07366.pdf)。
· 对新型可逆密度模型进行流程规格化。
· 连续时间序列模型:连续定义的动力学模型可以自动包含在任意时间到达的资料。
根据这篇文章,除了用ODENet代替ResNet来实现计算机视觉外,还存在一些目前无法付诸现实的应用:
· 将复杂的常微分方程压缩为单个动力学建模神经网络。
· 将其应用于缺少时间步长的时间序列。
· 可逆流程规格化(超出此文章的范围)。
对于不足之处,请参考原论文。有了充足的理论后,现在来看看一些例项。

学习动力学系统
正如之前所展示的,微分方程广泛应用于描述复杂的连续过程。当然,在现实生活中,我们把它们看成是离散的过程,最重要的是,在时间步长 t_i 中,很多观察结果可能会被忽略。假设想用一个神经网络来模拟这样一个系统。在经典的序列建模范例中,将如何处理这种情况?可能会运用递回神经网络,而递回神经网络甚至不是为它而设计的。在这一部分中,将考察神经常微分方程将如何处理这些情况。
设定如下:
1. 定义常微分方程本身,建模为PyTorch nn.Module()
2. 确定一个简单的(或不是真正的)神经网络,该神经网络将在从h_t到h_{t+1}的两个后续动力学步骤之间进行动力学建模,或者在动力学系统中,对 x_t和 x_{t+1}进行建模。
3. 执行通过常微分方程求解器反向传播的优化过程,使实际动力学和动力学建模之间的差异最小化。
在下面的所有实验中都会伴有神经网络(这足以用两个变数来建立简单的函式模型):
self.net = nn.Sequential(
nn.Linear(2, 50),
nn.Tanh(),
nn.Linear(50, 2),
)
所有更深层次的例子都受到了这个数据库的启发(https://nbviewer.jupyter.org/github/urtrial/neural_ode/),并给出了详尽的解释。在接下来的几小节中将展示所建立的动力学系统模型在程式码中的体现和系统如何随着时间演化以及ODENet如何拟合相图。
简单螺旋形函式
在此处以及所有后续的影象中,虚线代表拟合模型。
true_A = torch.tensor([[-0.1, 2.0], [-2.0, -0.1]])
class Lambda(nn.Module):
def forward(self, t, y):
return torch.mm(y, true_A)


上边是相空间,下边是时间-空间。直线代表真实的轨迹,虚线代表神经常微分方程系统学习的变化过程。
随机矩阵函式
true_A = torch.randn(2, 2)/2.


上边是相空间,下边是时间-空间。直线代表真实的轨迹,虚线代表神经常微分方程系统学习的变化过程。
Volterra-Lotka系统
a, b, c, d = 1.5, 1.0, 3.0, 1.0true_A = torch.tensor([[0., -b*c/d], [d*a/b, 0.]])


上边是相空间,下边是时间-空间。直线代表真实的轨迹,虚线代表神经常微分方程系统学习的变化过程。
非线性函式
true_A2 = torch.tensor([[-0.1, -0.5], [0.5, -0.1]])
true_B2 = torch.tensor([[0.2, 1.], [-1, 0.2]])
class Lambda2(nn.Module):
def __init__(self, A, B):
super(Lambda2, self).__init__()
self.A = nn.Linear(2, 2, bias=False)
self.A.weight = nn.Parameter(A)
self.B = nn.Linear(2, 2, bias=False)
self.B.weight = nn.Parameter(B)
def forward(self, t, y):
xTx0 = torch.sum(y * true_y0, dim=1)
dxdt = torch.sigmoid(xTx0) * self.A(y - true_y0) + torch.sigmoid(-xTx0) * self.B(y + true_y0)
return dxdt


上边是相空间,下边是时间-空间。直线代表真实的轨迹,虚线代表神经常微分方程系统学习的变化过程。
如上所示,单个“残差块”不能很好地学习这个过程,所以可能会使其更加复杂,以应对随后的函式。
神经网络函式
使用带有随机初始权值的多层感知器充分确定函式的引数:
true_y0 = torch.tensor([[1., 1.]])
t = torch.linspace(-15., 15., data_size)
class Lambda3(nn.Module):
def __init__(self):
super(Lambda3, self).__init__()
self.fc1 = nn.Linear(2, 25, bias = False)
self.fc2 = nn.Linear(25, 50, bias = False)
self.fc3 = nn.Linear(50, 10, bias = False)
self.fc4 = nn.Linear(10, 2, bias = False)
self.relu = nn.ELU(inplace=True)
def forward(self, t, y):
x = self.relu(self.fc1(y * t))
x = self.relu(self.fc2(x))
x = self.relu(self.fc3(x))
x = self.relu(self.fc4(x))
return x


左边是相空间,右边是时间-空间。直线代表真实的轨迹,虚线代表神经常微分方程系统学习的变化过程。
此处的2-50-2网络由于结构过于简单而严重失败,需要增加它的深度:
self.net = nn.Sequential(
nn.Linear(2, 150),
nn.Tanh(),
nn.Linear(150,50),
nn.Tanh(),
nn.Linear(50, 50),
nn.Tanh(),
nn.Linear(50, 2),
)


左边是相空间,右边是时间-空间。直线代表真实的轨迹,虚线代表神经常微分方程系统学习的变化过程。
现在基本上可以按照预期工作了,请勿忘记检查程式码。

作为生成式模型的神经常微分方程
论文作者还声称可以通过利用神经节点作为VAE框架的一部分来建立生成时间序列模型。那么工作原理如何呢?
图片来自原论文
· 首先,使用一些“标准”时间序列算法对输入序列进行编码,例如RNN,以进行过程中的初次嵌入。
· 通过神经常微分方程执行嵌入,实现“连续”嵌入
· 从VAE的“连续”嵌入中恢复初始序列。
为证明此概念,我们可以从储存库中重新执行程式码,它似乎在学习螺旋轨迹方面执行良好:

点代表取样噪声轨迹,蓝线是真实轨迹,橙线代表恢复轨迹和插值轨迹。
随后,我们可以从心电图(ECG)转换为 x(t)以为时间-空间,x`(t)为微分-空间的相图(如本文所示),并尝试拟合不同的VAE设定。这种用例对于像MAWI BandMawi Band这样的可穿戴装置可能非常有用,由于有噪声或中断的讯号,我们必须恢复它(实际上我们是在深度学习的帮助下完成的(the help of deep learning),但心电图是一个连续的讯号,不是吗?)不幸的是,它并没有很好地保持一致,显示出了所有过度拟合于单一形式的跳动迹象。


相空间。蓝线—实际轨迹,橙线—抽样轨迹以及噪声轨迹,绿线—自动编码轨迹。


时空。蓝线-实际讯号,橙线-抽样讯号和噪声讯号,绿线-自动编码讯号。
我们还可以尝试另一个实验:只在每个跳动的部分研究该自动编码器,并从中恢复整个波形(即外推出一个讯号)。不幸的是,无论如何进行超引数和资料预处理,将这段讯号向左或向右外推,都无法得出任何有意义的结果。也许读者们可以帮助解决这个问题。

下一步是什么?
很明显,神经常微分方程是为了研究相对简单的过程而设计的(这就是为什么标题中有“常”字),所以需要一个能够对更丰富的函式族建模的模型。已经有两种有趣的方法:
· 扩大神经常积分方程: https://github.com/EmilienDupont/augmented-neural-odes
· 神经跳动随机DE: https://www.groundai.com/project/neural-jump-stochastic-differential-equations/1
当然,我们还需要时间来进行探究。毕竟神经常微分方程现阶段无法投入实践。这本身是个伟大的想法,目前只存在两个实际应用:
· 在经典神经网络中使用ODESolve()层来平衡速度/精度折衷。
· 将常规常微分方程“压缩”到神经架构中,以将其嵌入标准资料科学流水线中。

留言 点赞 关注
我们一起分享AI学习与发展的干货
欢迎关注全平台AI垂类自媒体 “读芯术”