快速优化视角:基于张量与 SVM 技巧重构 LLM 中的单层注意力并以其矩阵乘法时间求解
数据结构与算法
2023-09-15 v1 机器学习
机器学习
摘要
大语言模型(LLMs)在变革我们日常生活的诸多方面发挥了关键作用。求解注意力回归是优化 LLM 的一项基本任务。在本工作中,我们致力于为单层注意力网络目标函数 L ( X , Y ) = ∑ j 0 = 1 n ∑ i 0 = 1 d ( ⟨ ⟨ exp ( A j 0 x ) , 1 n ⟩ − 1 exp ( A j 0 x ) , A 3 Y ∗ , i 0 ⟩ − b j 0 , i 0 ) 2 L(X,Y) = \sum_{j_0 = 1}^n \sum_{i_0 = 1}^d ( \langle \langle \exp( \mathsf{A}_{j_0} x ) , {\bf 1}_n \rangle^{-1} \exp( \mathsf{A}_{j_0} x ), A_{3} Y_{*,i_0} \rangle - b_{j_0,i_0} )^2 L ( X , Y ) = ∑ j 0 = 1 n ∑ i 0 = 1 d (⟨⟨ exp ( A j 0 x ) , 1 n ⟩ − 1 exp ( A j 0 x ) , A 3 Y ∗ , i 0 ⟩ − b j 0 , i 0 ) 2 提供可证明的保证。此处 A ∈ R n 2 × d 2 \mathsf{A} \in \mathbb{R}^{n^2 \times d^2} A ∈ R n 2 × d 2 是 A 1 ∈ R n × d A_1 \in \mathbb{R}^{n \times d} A 1 ∈ R n × d 与 A 2 ∈ R n × d A_2 \in \mathbb{R}^{n \times d} A 2 ∈ R n × d 间的克罗内克积。A 3 A_3 A 3 是 R n × d \mathbb{R}^{n \times d} R n × d 中的矩阵,A j 0 ∈ R n × d 2 \mathsf{A}_{j_0} \in \mathbb{R}^{n \times d^2} A j 0 ∈ R n × d 2 是 A \mathsf{A} A 的第 j 0 j_0 j 0 块。X , Y ∈ R d × d X, Y \in \mathbb{R}^{d \times d} X , Y ∈ R d × d 为我们欲学习的变量。B ∈ R n × d B \in \mathbb{R}^{n \times d} B ∈ R n × d 且 b j 0 , i 0 ∈ R b_{j_0,i_0} \in \mathbb{R} b j 0 , i 0 ∈ R 是 B B B 的第 j 0 j_0 j 0 行第 i 0 i_0 i 0 列元素,Y ∗ , i 0 ∈ R d Y_{*,i_0} \in \mathbb{R}^d Y ∗ , i 0 ∈ R d 是 Y Y Y 的第 i 0 i_0 i 0 列向量,x ∈ R d 2 x \in \mathbb{R}^{d^2} x ∈ R d 2 是 X X X 的向量化。在多层的 LLM 网络中,矩阵 B ∈ R n × d B \in \mathbb{R}^{n \times d} B ∈ R n × d 可视为某层的输出,A 1 = A 2 = A 3 ∈ R n × d A_1= A_2 = A_3 \in \mathbb{R}^{n \times d} A 1 = A 2 = A 3 ∈ R n × d 可视为某层的输入。x x x 的矩阵形式可视为 Q K ⊤ QK^\top Q K ⊤ ,Y Y Y 可视为 V V V 。我们提供一种迭代贪婪算法,以 O ~ ( ( T m a t ( n , n , d ) + T m a t ( n , d , d ) + d 2 ω ) log ( 1 / ϵ ) ) \widetilde{O}( ({\cal T}_{\mathrm{mat}}(n,n,d) + {\cal T}_{\mathrm{mat}}(n,d,d) + d^{2\omega}) \log(1/\epsilon) ) O (( T mat ( n , n , d ) + T mat ( n , d , d ) + d 2 ω ) log ( 1/ ϵ )) 时间将损失函数 L ( X , Y ) L(X,Y) L ( X , Y ) 训练至 ϵ \epsilon ϵ 误差内。此处 T m a t ( a , b , c ) {\cal T}_{\mathrm{mat}}(a,b,c) T mat ( a , b , c ) 表示将一个 a × b a \times b a × b 矩阵与另一个 b × c b \times c b × c 矩阵相乘的时间,ω ≈ 2.37 \omega\approx 2.37 ω ≈ 2.37 表示矩阵乘法的指数。
引用
@article{arxiv.2309.07418,
title = {A Fast Optimization View: Reformulating Single Layer Attention in LLM Based on Tensor and SVM Trick, and Solving It in Matrix Multiplication Time},
author = {Yeqi Gao and Zhao Song and Weixin Wang and Junze Yin},
journal= {arXiv preprint arXiv:2309.07418},
year = {2023}
}