AI 中文总结
该研究提出无需子空间假设的多头注意力学习算法,通过特定查询设计恢复规范头,可精确恢复参数且适用于近似输出,还扩展至单层 Transformer 的学习。
AI 中文摘要
我们研究基于黑盒输入输出访问学习多头softmax注意力的问题,学习者可查询任意实值 token 序列,仅观测最终 token 的标量输出。近期工作提出用 $O(d^2)$ 次值查询恢复单头参数 $(W,v)$;对于多头,该工作假设各头占据两两正交子空间以保证可识别性,若要分别对各头应用单头恢复算法,还需已知这些子空间的基。我们无需子空间假设,通过合并具有相同 $W_h$ 的头、求和对应 $v_h$ 并在和为零时丢弃合并头,从而恢复规范表示。通过改变 token 的副本数,我们的算法获得有理函数样本,其插值可分离规范头;添加选定 token 向量形成的额外查询则匹配不同查询中的同一头。当预言机输出及后续所有计算均精确时,学习者随机选择查询向量,以概率 1 恢复规范对 $\{(W_h,v_h):h\in[H]\}$(不计排列顺序);若 $H$ 已知,算法恰使用 $4Hd^2-2H+1$ 次最大长度为 $2H+1$ 的值查询;若仅已知 $H$ 的上界 $H_0$,则使用 $4H_0d^2-2H_0+1$ 次最大长度为 $2H_0+1$ 的值查询。对于近似预言机输出,我们给出参数误差不超过模型及查询依赖常数倍输出误差的条件。最后,我们将结果扩展至包含多头注意力及无偏置 ReLU 前馈网络的单层 Transformer,在额外条件下,无需单独学习前馈网络的算法即可恢复功能等价的 Transformer。
英文摘要
We study the problem of learning multi-head softmax attention from black-box input-output access. The learner may query arbitrary real-valued token sequences and observe only the scalar output at the final token. Recent work gives an algorithm using $O(d^2)$ value queries to recover the single-head parameters $(W,v)$. For multiple heads, the same work establishes identifiability under the assumption that the heads occupy pairwise orthogonal subspaces. Applying the single-head recovery algorithm separately to the heads additionally requires bases for these subspaces to be known. We recover a canonical representation by merging heads with the same $W_h$, summing their corresponding $v_h$, and discarding a merged head when this sum is zero, without these subspace assumptions. By varying the number of copies of a token, our algorithm obtains samples of a rational function whose interpolation separates the canonical heads. Additional queries formed by adding selected token vectors then match the same head across different queries. When the oracle outputs and all subsequent computations are exact, the learner chooses its query vectors at random and recovers the canonical pairs $\{(W_h,v_h):h\in[H]\}$ up to permutation with probability one. When $H$ is known, it uses exactly $4Hd^2-2H+1$ value queries of maximum length $2H+1$. If only a known upper bound $H_0$ is available, the algorithm uses $4H_0d^2-2H_0+1$ value queries of maximum length $2H_0+1$. For approximate oracle outputs, we give conditions under which the parameter error is at most a model- and query-dependent constant multiple of the output error. Finally, we extend our result to a one-layer Transformer with multi-head attention followed by a bias-free ReLU feed-forward network. Under additional conditions, we recover a functionally equivalent Transformer without relying on a separate algorithm for learning the feed-forward network.
Comments39 pages