使用线性自注意力Transformer对简单线性回归任务的闭式解进行上下文学习

arXiv cs.LG 论文

摘要

本文构造了一个具有线性自注意力的Transformer,该Transformer对简单线性回归执行闭式最小二乘解的上下文学习,利用层归一化来近似解析解,而非梯度下降。

arXiv:2607.15819v1 Announce Type: new 摘要:上下文学习是Transformer的一个显著特性,近年来引起了广泛关注。在许多上下文学习的研究中,已有研究表明Transformer能够实现线性和非线性回归问题的求解器,其中大多数实现使用梯度下降算法。然而,尚不清楚这些实现是否真正通过训练获得。本文构造了一个具有线性自注意力的Transformer,它在简单回归任务中以上下文方式学习最小二乘估计。关键在于,我们利用层归一化近似得到闭式(解析)解,而不是基于梯度下降算法的近似解。然后,我们展示了一个实验示例,其中我们的实现主要用于训练有l1正则化的Transformer,目标输出是最小二乘估计。
查看原文
查看缓存全文

缓存时间: 2026/07/20 09:31

# 使用线性自注意力Transformer进行简单线性回归任务的闭式解上下文学习
来源:https://arxiv.org/html/2607.15819
\(三重大学教育学院 1577 Kurima\-Machiya\-cho, Tsu, 514\-8507, Japan hagi@edu\.mie\-u\.ac\.jp\)

###### 摘要

上下文学习是Transformer的一项显著特性,近期受到了广泛关注。在许多关于上下文学习的研究中,已经证明Transformer能够实现线性和非线性回归问题的求解器,其中大多数实现了梯度下降算法。然而,这些实现是否真的通过训练获得仍不清楚。本文构建了一个具有线性自注意力的Transformer,该模型在简单回归任务中通过上下文学习最小二乘估计。关键点在于,闭式(解析)解是通过层归一化近似获得的,而非基于梯度下降算法的近似解。随后,我们展示了一个实验示例,其中当目标输出为最小二乘估计时,我们的实现在使用 \(\ell_1\) 正则化训练的Transformer中占主导地位。

关键词:上下文学习,线性自注意力,简单线性回归,层归一化

## 1 引言

上下文学习是Transformer的一项显著特性,而Transformer构成了GPT-3等大型语言模型的基础\[3 (https://arxiv.org/html/2607.15819#bib.bib3)\],并已成为近期研究的焦点。通过上下文学习,给定包含任务示例和新查询输入的提示,训练后的语言模型能够一次性为新查询生成相应的输出。人们自然地认为,Transformer通过训练获得了通过上下文学习解决任务的算法。在这方面,已有研究表明Transformer能够实现多种算法,特别是在回归任务中\[5 (https://arxiv.org/html/2607.15819#bib.bib5),11 (https://arxiv.org/html/2607.15819#bib.bib11),1 (https://arxiv.org/html/2607.15819#bib.bib1),4 (https://arxiv.org/html/2607.15819#bib.bib4),2 (https://arxiv.org/html/2607.15819#bib.bib2)\]。

文献\[5 (https://arxiv.org/html/2607.15819#bib.bib5)\] 实证研究了Transformer对机器学习中各类函数类(包括线性函数类)的上下文学习能力。特别地,对于线性函数,训练后的Transformer表现与最小二乘解相似。文献\[11 (https://arxiv.org/html/2607.15819#bib.bib11)\] 提供了线性自注意力层的显式构造,该层在均方误差损失上实现了梯度下降算法的单步迭代。此外,他们实证表明,多个自注意力层可以迭代地进行曲率校正,从而优于普通梯度下降算法。文献\[1 (https://arxiv.org/html/2607.15819#bib.bib1)] 证明了Transformer能够实现梯度下降算法以及岭回归的闭式解。文献\[4 (https://arxiv.org/html/2607.15819#bib.bib4)] 也指出了线性版本注意力与梯度下降算法之间的对应关系,声称Transformer执行隐式微调。他们还实证研究了上下文学习与显式微调之间的相似性。虽然文献\[5 (https://arxiv.org/html/2607.15819#bib.bib5),11 (https://arxiv.org/html/2607.15819#bib.bib11),1 (https://arxiv.org/html/2607.15819#bib.bib1),10 (https://arxiv.org/html/2607.15819#bib.bib10),4 (https://arxiv.org/html/2607.15819#bib.bib4)] 未考虑训练阶段,但文献\[13 (https://arxiv.org/html/2607.15819#bib.bib13)] 研究了简化Transformer架构中梯度流的学习动态,此时训练提示由线性回归数据集的随机实例组成,结论表明梯度流训练的Transformer通过上下文学习了一类线性函数。近期,文献\[2 (https://arxiv.org/html/2607.15819#bib.bib2)] 显示Transformer可以在上下文中实现广泛的标准机器学习算法,如最小二乘、岭回归和Lasso。与文献\[1 (https://arxiv.org/html/2607.15819#bib.bib1)] 相比,文献\[2 (https://arxiv.org/html/2607.15819#bib.bib2)] 精确评估了网络规模下的预测性能,并展示了接近最优的预测能力。文献\[2 (https://arxiv.org/html/2607.15819#bib.bib2)] 还展示了Transformer的算法选择能力,例如根据岭回归的验证误差选择正则化。然而,在这些工作中,这些算法是否真的通过Transformer的训练过程获得仍不清楚。本文针对一个简单回归任务,首先根据文献\[1 (https://arxiv.org/html/2607.15819#bib.bib1)] 构建了一个带线性自注意力的Transformer,实现了最小二乘估计的近似闭式(解析)解。然后我们给出一个数值示例,其中当目标为最小二乘估计时,该实现在使用 \(\ell_1\) 正则化训练的Transformer中占主导地位。

先前的工作\[5 (https://arxiv.org/html/2607.15819#bib.bib5),11 (https://arxiv.org/html/2607.15819#bib.bib11),1 (https://arxiv.org/html/2607.15819#bib.bib1),4 (https://arxiv.org/html/2607.15819#bib.bib4),2 (https://arxiv.org/html/2607.15819#bib.bib2)\] 研究了回归问题,发现Transformer实现的求解器基本上是梯度下降算法。这很自然,因为Transformer中的注意力机制计算输入的乘积,而这正是梯度下降法所要求的。然而,在线性回归问题中,最小二乘估计是通过计算矩阵逆来解析获得的,这需要除法。因此,梯度下降算法通过重复执行乘法和加法来隐式地计算这个除法。在这些工作中,文献\[1 (https://arxiv.org/html/2607.15819#bib.bib1)] 表明Transformer可以实现岭回归问题的闭式解,其中除法是通过层归一化实现的。我们对简单回归问题的最小二乘估计的构造正是基于这一见解。

我们设置中的Transformer由输入的线性变换、一组Transformer块堆叠、一个展平层和一个输出线性变换组成,其中Transformer块由带层归一化的多头线性自注意力块以及紧随其后的跳跃连接组成。因此,层归一化应用于线性自注意力块的输出与跳跃连接之和,然后再将层归一化后的输出应用于跳跃连接。该Transformer接收上下文样本和查询输入(预测点)作为提示,并输出针对查询输入的最小二乘估计。具体来说,我们构造中的层数、头数和模型维度分别为2、2和4,这非常小。简单回归问题的最小二乘估计的闭式解需要除以上下文样本中输入数据的方差。我们提供了一个具体的构造,在自然的输入形式下近似计算此闭式解,其中根据文献\[1 (https://arxiv.org/html/2607.15819#bib.bib1)\] 使用层归一化来执行除法。在展示此构造之后,我们进行数值实验,以表明当目标输出为最小二乘估计时,我们的实现在使用 \(\ell_1\) 正则化训练的Transformer中确实被使用。换句话说,通过训练,Transformer主要获得的是闭式解的计算,而非梯度下降算法的步骤。

本文组织如下。第2节阐述我们的问题设置。第3节给出基于上下文样本计算最小二乘估计的Transformer构造。第4节提供训练的数值示例,以证明该构造实际上有效。最后,第5节总结全文并讨论未来工作。

## 2 问题设置

### 2.1 记号

在本文中,\(\mathbf{O}_{I,J}\) 是 \(I \times J\) 零矩阵。注意,此记号也用于向量,此时 \(I\) 或 \(J\) 等于1。对于 \(I \times J\) 矩阵 \(\mathbf{A}\),\(\mathbf{A}[i,j]\) 是 \(\mathbf{A}\) 的 \((i,j)\) 元素,\(\mathbf{A}[i,:]\) 是 \(\mathbf{A}\) 的第 \(i\) 行向量,\(\mathbf{A}[:,j]\) 是 \(\mathbf{A}\) 的第 \(j\) 列向量,其中 \(i=1,\ldots,I\) 且 \(j=1,\ldots,J\)。当 \(\mathbf{A}\) 是 \(I \times 1\) 向量时,\(\mathbf{A}[i]\) 表示其第 \(i\) 个元素。

### 2.2 Transformer的上下文学习设置

接下来我们解释基于Transformer的简单线性回归任务的上下文学习。令 \(M\) 为Transformer的训练数据数量。令 \((x,y)\) 为简单回归问题中的一对输入-输出变量。对于每个 \(m \in \{1,\ldots,M\}\),生成一组 \(N\) 个 \((x,y)\) 的随机样本,记为 \(\{(x_{m,n},y_{m,n}): n=1,\ldots,N\}\)。我们将第 \(m\) 个提示(Transformer的输入)记为 \(\mathbf{P}_m\),它是一个 \((N+1) \times 3\) 矩阵,其第 \(n\) 行为:

\[
\mathbf{P}_m[n,:] = \begin{bmatrix} 1 & x_{m,n} & y_{m,n} \end{bmatrix},
\tag{1}
\]

其中我们定义 \(y_{m,N+1}:=0\) 且 \(x_{m,N+1}:=u_m\),即一个预测点;例如参见文献\[1 (https://arxiv.org/html/2607.15819#bib.bib1)\]。因此,输入序列长度为 \(N+1\),上下文样本数为 \(N\)。我们假设对于每个 \(x_{m,n}\),\(y_{m,n}\) 由下式生成:

\[
y_{m,n} = \theta_{m,0} + \theta_{m,1} x_{m,n} + \varepsilon_{m,n},
\tag{2}
\]

其中 \(\varepsilon_{m,1},\ldots,\varepsilon_{m,N}\),\(m=1,\ldots,M\) 是独立同分布的加性噪声,来自均值为0、方差 \(\sigma^2 < \infty\) 的概率分布。我们的目标是针对提示 \(\mathbf{P}_m\) 在 \(x=u_m\) 处获得最小二乘预测。因此,对于每个 \(m\),我们需要使用

\[
\mathbf{D}_m = \{(x_{m,n},y_{m,n}): n=1,\ldots,N\},
\tag{3}
\]

来计算最小二乘解,这是上下文样本的集合。关键在于,回归线对于每个 \(m\) 可以是不同的。因此,我们可以假设 \((\theta_{m,0},\theta_{m,1})\) 对于每个 \(m\) 从一个概率分布中采样。然而,在Transformer的构造中,我们对 \((\theta_{m,0},\theta_{m,1})\) 的潜在概率分布不做任何具体假设。虽然我们对 \(x_{m,n}\) 和 \(\varepsilon_{m,n}\) 的潜在概率分布也不做任何具体假设,但稍后我们将对上下文样本做出一个假设。

### 2.3 Transformer的训练数据

对于第 \(m\) 个训练数据 \(\mathbf{D}_m\)(定义于 (3)),我们定义

\[
\overline{x}_m := \frac{1}{N} \sum_{n=1}^{N} x_{m,n}
\tag{4}
\]
\[
\overline{y}_m := \frac{1}{N} \sum_{n=1}^{N} y_{m,n}
\tag{5}
\]
\[
V_m := \frac{1}{N} \sum_{n=1}^{N} (x_{m,n} - \overline{x}_m)^2
\tag{6}
\]
\[
C_m := \frac{1}{N} \sum_{n=1}^{N} (x_{m,n} - \overline{x}_m) (y_{m,n} - \overline{y}_m).
\tag{7}
\]

对于简单回归问题,容易看出基于 \(\mathbf{D}_m\) 的最小二乘解在 \(x=u_m\) 处的预测由下式给出:

\[
\widehat{y}_m(u_m) := \overline{y}_m + \frac{C_m}{V_m} (u_m - \overline{x}_m)
\tag{8}
\]
(例如参见文献\[14 (https://arxiv.org/html/2607.15819#bib.bib14)\])。

下面,我们构建一个Transformer,它接收 \(\mathbf{P}_m\) 作为输入,并对每个 \(m=1,\ldots,M\) 输出 \(\widehat{y}_m(u_m)\)。因此,该Transformer在每个 \(m\) 处使用上下文样本 \(\mathbf{D}_m\) 计算预测点 \(u_m\) 处的最小二乘估计。

### 2.4 线性自注意力

我们定义线性自注意力(LSA),它接收一个 \((N+1) \times D\) 矩阵 \(\mathbf{Q}\) 作为输入,并输出一个 \((N+1) \times D\) 矩阵,定义为:

\[
{\rm LSA}_{\boldsymbol{\Theta}}(\mathbf{Q}) := (\mathbf{Q} \mathbf{W}_3) ((\mathbf{Q} \mathbf{W}_1^\top)^\top (\mathbf{Q} \mathbf{W}_2)) = (\mathbf{Q} \mathbf{W}_3) (\mathbf{W}_1 \mathbf{Q}^\top \mathbf{Q} \mathbf{W}_2),
\tag{9}
\]

其中 \(\boldsymbol{\Theta} = \{\mathbf{W}_1, \mathbf{W}_2, \mathbf{W}_3\}\) 是一个参数的有序集合,且 \(\mathbf{W}_1, \mathbf{W}_2\) 和 \(\mathbf{W}_3\) 是 \(D \times D\) 矩阵。此LSA的操作在文献\[6 (https://arxiv.org/html/2607.15819#bib.bib6)\] 中有详细讨论。

### 2.5 Transformer结构

本文考虑的Transformer结构如图1所示,其中灰色块为操作块。此处的Transformer由输入线性变换、堆叠的Transformer块、展平层和输出线性变换组成(图1(a))。Transformer块由多头LSA(MHLSA)块以及紧随其后的带层归一化的跳跃连接组成(图1(b))。它接收第 \(m\) 个提示 \(\mathbf{P}_m\)(定义于 (1)),为简单起见,输出一个标量值,其目标为预测点 \(u_m\) 处的最小二乘估计,即由 (8) 给出的 \(\widehat{y}_m(u_m)\)。

参照图注

(a) 主流程

参照图注

(b) Transformer块

图1: Transformer的结构

令 \(\mathbf{W}_{\rm in}\) 和 \(\mathbf{b}_{\rm in}\) 分别为 \(D \times 3\) 输入权重矩阵和 \(D \times 1\) 输入偏置向量。提示 \(\mathbf{P}_m\) 通过仿射变换 \((\mathbf{W}_{\rm in}, \mathbf{b}_{\rm in})\) 的嵌入记为 \(\mathbf{Q}_{m,1}\),其大小为 \((N+1) \times D\)。更精确地,我们定义

\[
\mathbf{Q}_{m,1}[n,:]^\top := \mathbf{W}_{\rm in} \mathbf{P}_m[n,:]^\top + \mathbf{b}_{\rm in}
\tag{10}
\]

对于 \(n=1,\ldots,N+1\)。层数记为 \(L\)。对于第 \(m\) 个训练数据...

相似文章

Exact Linear Attention

arXiv cs.LG

本文介绍了一种名为Exact Linear Attention (ELA) 的机制,该机制通过利用核函数分解,在不引入近似误差的情况下实现了Transformer注意力的线性计算复杂度,并通过约束核函数解决了梯度爆炸和词元稀释问题。文中还提出了包括超链接(Hyper Link)、记忆叶(Memory Lobe)以及面向混合专家模型的路由偏置在内的工程创新。

LLT: 用于PDE算子学习的局部线性Transformer

arXiv cs.LG

介绍LLT,一种基于Transformer的神经算子,它将线性全局注意力与局部空间混合相结合,用于PDE学习。在多个PDE问题上,与基线方法相比,它实现了具有竞争性的精度和更快的训练速度。

变分线性注意力:用于长上下文 Transformer 的稳定联想记忆

arXiv cs.LG

本文介绍了变分线性注意力(VLA),这是一种用于稳定长上下文 Transformer 中线性注意力机制记忆状态的方法。VLA 将记忆更新重构为在线正则化最小二乘问题,证明了状态范数的有界性,并展示了相较于标准线性注意力和 DeltaNet 显著的速度提升以及更高的检索准确性。