Skip to main content
QUICK REVIEW

[论文解读] Transformers learn to implement preconditioned gradient descent for in-context learning

Kwangjun Ahn, Xiang Cheng|arXiv (Cornell University)|Jun 1, 2023
Stochastic Gradient Optimization Techniques被引用 6
一句话总结

该论文通过分析损失曲面,表明在随机线性回归实例上训练的Transformer模型学会通过实现预条件梯度下降来学习。对于单层注意力机制,全局最优解对应于一个适应数据分布和数据不足引起的方差的预条件梯度下降的一步;更深的Transformer则实现多步迭代,其临界点与自适应优化算法(包括GD++)相匹配。

ABSTRACT

Several recent works demonstrate that transformers can implement algorithms like gradient descent. By a careful construction of weights, these works show that multiple layers of transformers are expressive enough to simulate iterations of gradient descent. Going beyond the question of expressivity, we ask: Can transformers learn to implement such algorithms by training over random problem instances? To our knowledge, we make the first theoretical progress on this question via an analysis of the loss landscape for linear transformers trained over random instances of linear regression. For a single attention layer, we prove the global minimum of the training objective implements a single iteration of preconditioned gradient descent. Notably, the preconditioning matrix not only adapts to the input distribution but also to the variance induced by data inadequacy. For a transformer with $L$ attention layers, we prove certain critical points of the training objective implement $L$ iterations of preconditioned gradient descent. Our results call for future theoretical studies on learning algorithms by training transformers.

研究动机与目标

  • 研究Transformer是否可以通过在随机问题实例上训练来学习基于梯度的优化算法,而非依赖手工设计的权重。
  • 分析在随机线性回归实例上训练的线性Transformer的损失曲面,以理解优化算法如何从训练中涌现。
  • 表征Transformer参数空间中全局最小值和临界点的结构,以识别其所实现的优化算法。
  • 弥合Transformer理论表达能力与其在上下文学习中实际行为之间的差距,特别是针对基于梯度的方法。
  • 通过实验验证理论发现,表明学习到的临界点与已知的自适应优化算法(如GD++)相匹配。

提出的方法

  • 使用非Softmax注意力机制,分析在随机各向同性线性回归实例上训练的单层线性Transformer的损失曲面。
  • 证明训练目标的全局最小值对应于一个预条件梯度下降的单步迭代,其中预条件矩阵同时适应输入数据分布和由数据不足引起的方差。
  • 在参数空间中引入稀疏性条件,以限制搜索范围至k步自适应梯度优化算法,从而实现对多层Transformer临界点的表征。
  • 对于更深的Transformer(L层),证明某些临界点实现了具有数据相关预条件的L步预条件梯度下降。
  • 放宽稀疏性条件以研究完整参数空间,揭示一个临界点对应于一种新型基于梯度的算法,该算法将梯度步骤与线性变换结合以进一步改善条件性。
  • 通过可视化学习到的权重并比较Transformer预测器与标准优化基线(GD、预条件GD、OLS)的测试损失,实验验证理论发现。
(a) $\operatorname{Dist}(\Sigma^{1/2}A_{0}\Sigma^{1/2},I)$
(a) $\operatorname{Dist}(\Sigma^{1/2}A_{0}\Sigma^{1/2},I)$

实验结果

研究问题

  • RQ1在随机线性回归实例上训练的Transformer是否能通过非凸优化学会实现梯度下降?
  • RQ2在随机回归问题上训练的单层线性Transformer的全局最小值实现的是哪种优化算法?
  • RQ3在更深的Transformer(L层)中,临界点如何对应于迭代优化算法?
  • RQ4当移除对参数的稀疏性约束时,学习到的算法会发生什么变化?
  • RQ5理论上推导出的临界点是否与训练后Transformer中经验观察到的行为一致?

主要发现

  • 单层线性Transformer的全局最小值实现了单步预条件梯度下降,其中预条件矩阵适应输入数据分布和由数据不足引起的方差。
  • 在稀疏性条件下,两层Transformer的全局最小值对应于自适应步长的梯度下降,有效实现了k步自适应优化过程。
  • 在更深的Transformer(L层)中,训练目标的某些临界点实现了具有数据相关预条件的L步预条件梯度下降。
  • 在无稀疏性约束时,一个临界点对应于一种新型算法,该算法将梯度下降与线性变换结合以进一步改善条件性,当数据协方差为各向同性时,其与GD++算法一致。
  • 实验验证表明,三层Transformer学习到的权重与理论上的驻点一致,且该点的目标函数值接近零,表明其为全局最优解。
  • 测试损失比较结果表明,Transformer学习到的预测器性能与三步预条件梯度下降相当,验证了理论结果与实际优化行为的一致性。
(b) $\operatorname{Dist}(\Sigma^{1/2}A_{1}\Sigma^{1/2},I)$
(b) $\operatorname{Dist}(\Sigma^{1/2}A_{1}\Sigma^{1/2},I)$

更好的研究,从现在开始

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

无需绑定信用卡

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