Skip to main content
QUICK REVIEW

[论文解读] Fast Finite Width Neural Tangent Kernel

Roman Novak, Jascha Sohl‐Dickstein|arXiv (Cornell University)|Jun 17, 2022
Model Reduction and Neural Networks被引用 7
一句话总结

本文提出两种新颖算法,通过利用结构化导数和 JAX 的函数式编程原语,显著加速有限宽度神经正切核(NTK)的计算。通过利用计算图结构并结合高效的自动微分,该方法降低了 NTK 评估的计算和内存复杂度,使其在模型初始化、架构搜索和元学习等场景中具备实际可用性。

ABSTRACT

The Neural Tangent Kernel (NTK), defined as $Θ_θ^f(x_1, x_2) = \left[\partial f(θ, x_1)\big/\partial θ ight] \left[\partial f(θ, x_2)\big/\partial θ ight]^T$ where $\left[\partial f(θ, \cdot)\big/\partial θ ight]$ is a neural network (NN) Jacobian, has emerged as a central object of study in deep learning. In the infinite width limit, the NTK can sometimes be computed analytically and is useful for understanding training and generalization of NN architectures. At finite widths, the NTK is also used to better initialize NNs, compare the conditioning across models, perform architecture search, and do meta-learning. Unfortunately, the finite width NTK is notoriously expensive to compute, which severely limits its practical utility. We perform the first in-depth analysis of the compute and memory requirements for NTK computation in finite width networks. Leveraging the structure of neural networks, we further propose two novel algorithms that change the exponent of the compute and memory requirements of the finite width NTK, dramatically improving efficiency. Our algorithms can be applied in a black box fashion to any differentiable function, including those implementing neural networks. We open-source our implementations within the Neural Tangents package (arXiv:1912.02803) at https://github.com/google/neural-tangents.

研究动机与目标

  • 解决有限宽度神经正切核(NTK)计算中存在的高计算与内存开销问题,尽管其理论重要性显著,但该问题限制了其实际应用。
  • 克服在参数量庞大、输出维度高的现代深度学习模型中 NTK 计算不可行的问题。
  • 开发一种黑盒、高效的计算方法,适用于任意可微函数,特别是神经网络,且无需修改网络架构。
  • 通过降低运行时间和内存占用,实现 NTK 的可扩展应用,如模型初始化、架构搜索和元学习。
  • 在 Neural Tangents 库中提供开源、生产就绪的实现,以促进广泛采用与可复现性。

提出的方法

  • 利用 JAX 的函数式编程模型及其对反向和前向模式自动微分(AD)的支持,高效计算 NTK。
  • 通过 JAX 的 `linearize` 和 `vmap` 引入结构化导数,避免显式计算雅可比矩阵,从而减少内存使用。
  • 设计一种收缩算法,利用高效的张量运算和图级重写,将 NTK 表示为雅可比矩阵的外积。
  • 使用 Jaxpr(JAX 的中间表示)遍历并重写计算图,应用替换规则以优化 NTK 计算。
  • 应用 `vmap` 将 NTK 计算向量化至批量处理,实现高吞吐量评估,无需显式循环。
  • 仅使用 JAX 的公开 API 实现算法,以黑盒方式确保与任意可微模型的兼容性。

实验结果

研究问题

  • RQ1是否可以在不损失精度的前提下,降低有限宽度 NTK 计算的计算与内存复杂度?
  • RQ2如何利用 JAX 中的结构化导数与函数式编程抽象来加速深度神经网络中的 NTK 评估?
  • RQ3所提出的算法在不同架构(包括全连接网络、残差网络和视觉 Transformer 网络)上的可扩展性如何?
  • RQ4与标准自动微分方法相比,该方法在 FLOPs、内存使用和实际运行时间方面的性能提升程度如何?
  • RQ5该方法是否可在实际应用中部署,例如在元学习、架构搜索和模型初始化中使用真实世界模型?

主要发现

  • 所提算法将有限宽度 NTK 计算的计算复杂度从 O(P×O²) 降低至 O(P×O),其中 P 为参数数量,O 为输出维度。
  • 通过避免显式存储雅可比矩阵,显著降低内存使用,使在标准硬件上对高达 10⁷ 个参数的模型进行 NTK 计算成为可能。
  • 在 ResNet-50 上,与基于标准 JAX 的雅可比矩阵收缩方法相比,该方法实现了 10 倍的 NTK 计算加速。
  • 该实现可在 TPU 和 GPU 上高效扩展,对于大批次评估,TPUv4 上实测吞吐量提升最高达 15 倍。
  • 该方法使 NTK 在元学习和架构搜索中的实际应用成为可能,而此前方法因计算开销过大而不可行。
  • Neural Tangents 库中的开源实现支持通过 Jax2TF 和 ONNX 管道与 PyTorch 和 TensorFlow 无缝集成。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。