Neural Tangent Kernel

Abstract. 【Under Construction】

简介

神经网络

考虑一个 $L + 1$ 层的神经网络,诸层依次有 $n_0, …, n_L$ 个神经元,并且配备一个 Lipschitz 连续、二阶可导的非线性激活函数 $\sigma : \mathbb{R}\rightarrow \mathbb{R}$。一个这种形状的神经网络有以下参数:

  • 线性映射. $\mathbf{W}^\ell \in \mathbb{R}^{n_\ell \times n_{\ell + 1}}$。
  • 偏置项. $\boldsymbol{b}^\ell \in \mathbb{R}^{n_{\ell + 1}}$。

所有参数统称为 $\theta\in \mathbb{R}^P$,其中总参数量 $P = \sum_{i = 0}^{L - 1} (n_i + 1)n_{i + 1}$。给定一组参数,这组参数被所谓的 $L$ 层神经网络实现函数 $F^L : \mathbb{R}^P\rightarrow (\mathbb{R}^{n_0}\rightarrow \mathbb{R}^{n_L})$ 映射到下面的 $\mathbb{R}^{n_0}\rightarrow \mathbb{R}^{n_L}$ 的函数 $f_\theta$:

  • 输入 $\boldsymbol{x}^0\in \mathbb{R}^{n_0}$;
  • $\widetilde{\boldsymbol{x}}^{i + 1} = \frac{1}{\sqrt{n_i}}\mathbf{W}^i\boldsymbol{x}^i + \beta \boldsymbol{b}^i$;
  • $\boldsymbol{x}^i = \sigma(\widetilde{\boldsymbol{x}}^i)$(逐元作用);
  • 输出 $\boldsymbol{x}^L$。

上方的线性层中因子 $1 / \sqrt{n_i}$ 的作用基本上是让网络在每个参数独立初始化为标准正态分布 $\mathcal{N}(0, 1)$ 和无穷宽假设之下不爆炸。如果不加这个缩放因子,等价于网络参数初始化为 $\mathcal{N}(1, 1 / n_i)$(LeCun Initialization),和大名鼎鼎的 Kaiming Initialization 差了一个和 $\sigma$ 有关的缩放因子 $C_\sigma$。

对于函数空间 $\mathcal{F} : \mathbb{R}^{n_0}\rightarrow \mathbb{R}^{n_L}$,任取输入概率分布 $p^{in}$,我们可以用如下内积度量这个函数空间的几何结构:

$$
\left\langle f, g\right\rangle_{p^{in}} = \mathbb{E}_{\boldsymbol{x}\in p^{in}} [f(\boldsymbol{x})^\top g(\boldsymbol{x})]
$$

在分析深度学习中的具体实例之时,$p^{in}$ 通常被取作某个离散的数据集 $\mathcal{D} = \{x_1, …, x_N\}$ 上的均匀分布:

$$
\rho = \frac{1}{N}\sum_{i = 1}^N \delta_{x_i}
$$

其中 $\delta$ 是狄拉克测度。

核梯度

在优化一个神经网络时,我们往往取一个凸的损失函数 $\mathcal{L} : \mathbb{R}^{n_L}\rightarrow \mathbb{R}$,比如均方误差中有

$$
\mathcal{L}(\boldsymbol{x}, \boldsymbol{x}^*) = \frac 12\left\lVert\boldsymbol{x} - \boldsymbol{x}^*\right\rVert^2
$$

以此对于给定的概率分布 $\rho$ 和 ground truth 函数,诱导出损失泛函 $C : \mathcal{F}\rightarrow \mathbb{R}$,定义为

$$
C(f) = \int \mathcal{L}(f(\boldsymbol{x}), f^*(\boldsymbol{x})) \rho(\boldsymbol{x})\mathrm{d}\boldsymbol{x} \label{eq:loss_functional}
$$

本质上,我们是在寻找最优的一个参数 $\theta^*$ 使得 $(C\circ F)(\theta^*)$ 是最小化。然而,实践中 $C\circ F : \mathbb{R}^P\rightarrow \mathbb{R}$ 一般不凸,这给你用如下动力学方程去优化网络参数所得结果的分析带来了诸多疑云:

$$
\frac{\mathrm{d}\theta}{\mathrm{d}t} = -\nabla_\theta (C\circ F) (\theta) \label{eq:finite_dataset}
$$

然而,因为 $C$ 在 $\mathcal{F}$ 上是凸的,我们不妨考虑在这个梯度流中,函数本身是如何变化的。有

$$
\begin{aligned}
\frac{\mathrm{d} f_\theta(\boldsymbol{x})}{\mathrm{d}t} &= \nabla_\theta f_\theta(\boldsymbol{x})^\top \frac{\mathrm{d}\theta}{\mathrm{d}t} \\
&= -\nabla_\theta f(\boldsymbol{x})^\top \nabla_{\theta} (C\circ F) (\theta) \\
&= -\nabla_\theta f(\boldsymbol{x})^\top \int \nabla_\theta f_\theta(\boldsymbol{z}) \frac{\partial \mathcal{L}}{\partial f_\theta(\boldsymbol{z})} \rho(\boldsymbol{z}) \mathrm{d}\boldsymbol{z} \\
&= -\int {\color{blue} \nabla_\theta f(\boldsymbol{x})^\top \nabla_\theta f(\boldsymbol{z})} { \frac{\partial \mathcal{L}}{\partial f_\theta(\boldsymbol{z})}\rho(\boldsymbol{z})}\mathrm{d}\boldsymbol{z}
\end{aligned} \label{eq:gradient_descent_func}
$$

Hint.

我们现在考虑有限的数据集 $(\ref{eq:finite_dataset})$,定义

$$
\begin{aligned}
\boldsymbol{u} &= \begin{pmatrix}
f(\boldsymbol{x}_1) \\
\vdots \\
f(\boldsymbol{x}_N)
\end{pmatrix}, \qquad
\frac{\partial \mathcal{L}}{\partial \boldsymbol{u}} = \frac{1}{N}\begin{pmatrix}
\frac{\partial \mathcal{L}}{\partial f_\theta(\boldsymbol{x}_1)} \\
\vdots \\
\frac{\partial \mathcal{L}}{\partial f_\theta(\boldsymbol{x}_N)}
\end{pmatrix}, \\
\mathbf{K} &= \begin{pmatrix}
\nabla_\theta f_\theta(\boldsymbol{x}_1)^\top \nabla_\theta f_\theta(\boldsymbol{x}_1) & \cdots & \nabla_\theta f_\theta(\boldsymbol{x}_1)^\top \nabla_\theta f_\theta(\boldsymbol{x}_N) \\
\vdots & & \vdots \\
\nabla_\theta f_\theta(\boldsymbol{x}_N)^\top \nabla_\theta f_\theta(\boldsymbol{x}_1) & \cdots & \nabla_\theta f_\theta(\boldsymbol{x}_1)^\top \nabla_\theta f_\theta(\boldsymbol{x}_N)
\end{pmatrix}
\end{aligned}
$$

可以整理出

$$
\frac{\mathrm{d}\boldsymbol{u}}{\mathrm{d}t} = -\mathbf{K}\frac{\partial \mathcal{L}}{\partial \boldsymbol{u}}
$$

显然在神经网络被过参数化($P \gg N$)的训练中,$\mathbf{K}$ 始终是一个正定矩阵。如果 $\mathbf{K}$(近似)保持不变,那么单看样本上的函数值,这个优化过程就是关于诱导范数 $\left\lVert\cdot\right\rVert_{\mathbf{K}^{-1}}$ 的最陡峭梯度下降(Steepest Descent Method)。而 $\mathcal{L}$ 关于 $\boldsymbol{u}$ 是凸的,因此函数将收敛于唯一的最低点。

从形式上看,这无非是 $\partial L / \partial f_\theta$ 乘上了一个线性变换。此处,我们补充一套记号来刻画这种线性变换及其正定性:

定义 2.1(高维核). 一个高维核 $K$ 是一个 $\mathcal{F}\otimes \mathcal{F}$ 上的二阶对称张量。即

$$
K : \mathbb{R}^{n_0}\times \mathbb{R}^{n_0} \rightarrow \mathbb{R}^{n_L\times n_L}
$$

且满足 $K(\boldsymbol{x}, \boldsymbol{y}) = K(\boldsymbol{y}, \boldsymbol{x})^\top$。

对于任意一个高维核 $K$,在概率分布 $\rho$ 之下,可以诱导出核内积

$$
\left\langle f, g\right\rangle_K = \mathbb{E}_{\boldsymbol{x}\sim \rho, \boldsymbol{z}\sim \rho} [f(\boldsymbol{x})^\top K(\boldsymbol{x}, \boldsymbol{z}) g(\boldsymbol{z})]
$$

如果 $K$ 还满足正定性:

$$
\forall f, \quad \left\langle f, f\right\rangle_K\geq 0 \quad \wedge \quad \left\lVert f\right\rVert_\rho > 0 \rightarrow \left\langle f, f\right\rangle > 0
$$

还可以进一步诱导出核范数 $\left\lVert\cdot\right\rVert_K$。

定义 2.2(变分导数). 对于一个式 $\ref{eq:loss_functional}$ 形式泛函 $C : \mathcal{F}\rightarrow \mathbb{R}$,其变分导数定义为如下线性泛函:

$$
\partial C|_f = \delta \mapsto \lim_{\varepsilon \rightarrow 0} \frac{C(f + \varepsilon \delta)}{\varepsilon}
$$

该泛函的线性性可以这样算出:对于任意式 $\ref{eq:loss_functional}$ 形式的泛函有

$$
\begin{aligned}
\partial C|_f &= \frac{\partial}{\partial \varepsilon}\int \mathcal{L}(f(\boldsymbol{x}) + \varepsilon \delta(\boldsymbol{x}), f^*(\boldsymbol{x})) \rho(\boldsymbol{x}) \mathrm{d}\boldsymbol{x} \\
&= \int \frac{\partial L}{\partial f(\boldsymbol{x})}^\top \delta(\boldsymbol{x})\rho(\boldsymbol{x})\mathrm{d}x = \left\langle \frac{\partial \mathcal{L}}{\partial f(\boldsymbol{x})}, \delta(\boldsymbol{x})\right\rangle
\end{aligned}
$$

因此,任意泛函的变分导数 $\partial C|_f$ 都具有 $\left\langle d_f, \cdot\right\rangle_\rho$ 的形式,也就是对偶空间 $\mathcal{F}^*$ 中的一个元素。现在我们来定义一个高维核怎么作用在 $\mathcal{F}^*$ 上。我们知道对于任意的 $K$,锁定一个列号和它的第二元得到 $K_{\cdot, i}(\boldsymbol{x}, \cdot)$,这也是一个 $\mathcal{F}$ 中的元素,因此不妨自然地定义 $\Phi_K : \mathcal{F}^*\rightarrow \mathcal{F}$,其中

$$
\Phi_K(d^*)_i = \boldsymbol{x}\mapsto \left\langle d, K_{\cdot, i}(\cdot, \boldsymbol{x})\right\rangle
$$

定义 2.3(核梯度). 原变分导数经过核 $K$ 作用之后得到的东西称为核梯度,记

$$
\nabla_K C|_f = \phi_K(\partial C|_f)
$$

神经正切核

将式 $\ref{eq:gradient_descent_func}$ 中的蓝色部分写成 $\mathcal{F}\otimes \mathcal{F}$ 的形式,当然你完全可以使用其他的记号,这里只是为了对齐原论文。

定义 3.1(神经正切核). 对于一个 $L$ 层的神经网络,参数 $\theta$ 处的神经正切核定义为

$$
\Theta^L(\theta) = \sum_{i = 1}^P \frac{\partial}{\partial \theta_i} F(\theta)\otimes \frac{\partial}{\partial \theta_i} F(\theta)
$$

对比式 $\ref{eq:gradient_descent_func}$,可以发现对原函数做梯度下降时,神经网络所表示的函数会在函数空间中关于神经正切核做核梯度下降,动力学方程如下:

$$
\frac{\mathrm{d}f_{\theta}}{\mathrm{d}t} = -\nabla_{\Theta(\theta)} C|_{f_\theta}
$$

本节主要证明以下结果:

  1. 初始化. 如果使用 LeCun 初始化,那么
    1. 在无穷宽极限下,网络依分布收敛于一个高斯过程;
    2. NTK 依概率收敛于某个确定的核;
  2. 训练时.
    1. 在无穷宽极限下,任意时刻的 NTK 一致收敛于上面那个确定的核;
    2. $L\geq 2$ 时 NTK 保持正定。

接下来我们逐一证明这些结果。

定理 3.1. 对于任意 $L$ 层的神经网络(定义见第 2 节),若其中参数被独立同分布地初始化为 $\mathcal{N}(0, 1)$,当 $n_1, …, n_{L - 1}$ 依次趋向 $\infty$ 时,有网络诸层输出的第 $k$ 维 $f_{\theta, k}(\boldsymbol{x})$($k = 1, …, n_L$)依分布收敛于独立同分布的的高斯过程,这些高斯过程的均值均为 $0$,协方差满足以下递推关系

$$
\begin{aligned}
\Sigma^1(\boldsymbol{x}, \boldsymbol{x}’) &= \frac{1}{n_0}\boldsymbol{x}^\top \boldsymbol{x}’ + \beta^2 \\
\Sigma^{L + 1}(\boldsymbol{x}, \boldsymbol{x}’) &= \mathbb{E}_{f\sim \mathcal{GP}(0, \Sigma^L)}[\sigma(f(\boldsymbol{x}))\sigma(f(\boldsymbol{x}’))] + \beta^2
\end{aligned}
$$

证明. 施归纳于 $L$。当 $L = 1$ 时,没有隐藏层,且

$$
f_\theta(\boldsymbol{x}) = \frac{1}{\sqrt{n_0}}\mathbf{W}^0\boldsymbol{x} + \beta \boldsymbol{b}^0
$$

很明显 $f_\theta$ 是高斯过程,均值和方差容易计算,分别就是 $0$ 和上面的 $\Sigma^1$。

而当 $L\geq 1$ 时,希望归纳到 $L + 1$。这时候考虑 $f_{\theta, k}(\boldsymbol{x})$ 是什么:

$$
f_{\theta, k}(\boldsymbol{x}) = \frac{1}{\sqrt{n_L}}\sum_{i = 1}^{n_L} W_{ki}^L \sigma(\boldsymbol{x}_i^{L}(\boldsymbol{x})) + \beta \boldsymbol{b}_k^L
$$

取输入 $\boldsymbol{x}_1, …, \boldsymbol{x}_n$,根据多元中心极限定理,在 $n_L$ 趋向无穷时,$[f_{\theta, k}(\boldsymbol{x}_1), …, f_{\theta, k}(\boldsymbol{x}_n)]$ 依分布收敛于高斯分布。简单计算一下两个变量时的协方差,容易发现确实就是要证明的形式。

$\blacksquare$

定理 3.2. 对于任意 $L$ 层的神经网络,若其中参数被独立同分布地初始化为 $\mathcal{N}(0, 1)$,当 $n_1, …, n_{L - 1}$ 依次趋向 $\infty$ 时,有 $\Theta^L(\theta)$ 依概率收敛于 $\Theta^L_\infty \mathbf{I}$。其中 $\Theta^L_\infty : \mathbb{R}^{n_0}\times \mathbb{R}^{n_0}\rightarrow \mathbb{R}$,服从如下递推关系:

$$
\begin{aligned}
\Theta_\infty^1(\boldsymbol{x}, \boldsymbol{x}’) &= \Sigma^1(\boldsymbol{x}, \boldsymbol{x}’) \\
\Theta_\infty^{L + 1}(\boldsymbol{x}, \boldsymbol{x}’) &= \Theta_\infty^L(\boldsymbol{x}, \boldsymbol{x}’)\dot{\Sigma}(\boldsymbol{x}, \boldsymbol{x}’) + \Sigma^{L + 1}(\boldsymbol{x}, \boldsymbol{x}’)
\end{aligned}
$$

其中 $\dot{\Sigma}(\boldsymbol{x}, \boldsymbol{x}’) = \mathbb{E}_{f\sim \mathcal{GP}(0, \Sigma^L)}(\dot{\sigma}(f(\boldsymbol{x}))\dot{\sigma}(f(\boldsymbol{x}’)))$。

证明. 依旧施归纳于 $L$。