arXiv:cs.LG(机器学习,全量分类)· Rub\'en Dar\'io Guerrero·· 15 小时前AI 评分35
不漂移的方向:Transformer 注意力机制的 Stiefel 流形路由
Directions That Don't Drift: Stiefel Manifold Routing for Transformer Attention
AI 导读
研究将注意力中的查询和键投影约束到 Stiefel 流形,并用带切空间投影、步长范数上限和极分解回缩的黎曼 Adam 优化。在 CIFAR-10 patch 基准(n=10k)上,该方法比 AdamW 提升 +6.79 个百分点(12/12 配对起点,t=38.33),增益随数据量从 n=1k 时的 +1.9pp 增至 n=50k 时的 +6.7pp。
正文
Abstract:The query and key projections $\WQ,\WK$ in attention are almost always trained by Euclidean optimizers with no geometric constraint. We constrain them to the Stiefel manifold and optimize with a Riemannian Adam carrying one scalar second moment per frame---the form of \citet{becigneul2019}, here extended to the compact, non-Hadamard $\St(d,r)$ with a tangent projector, step-norm cap, and polar retraction. Four propositions prove steepest descent in the embedded metric, gradient-scale independence, well-conditioning, and exact $\mathrm{O}(d)$-equivariance. A fifth records that weight decay has \emph{identically zero} Riemannian gradient on $\St(d,r)$ ($W{=}WI_r$ lies in the normal space), so decay cannot act on the constrained frames. On a CIFAR-10 patch benchmark at $n{=}10\mathrm{k}$ this rule gains $\mathbf{+6.79}$\,pp over AdamW across 12 paired starts ($t{=}38.33$, $12/12$); earlier fixed-step Riemannian SGD gains $+1.97$\,pp, of which $+1.69$\,pp comes from frozen orthonormal initialization alone. The corrected Adam's lead grows with data: $+1.9$\,pp at $n{=}1\mathrm{k}$ to $+6.7$\,pp at $n{=}50\mathrm{k}$. A 12-seed ablation credits all gain to the scale-free step ($+4.63$\,pp, $12/12$), nothing to the projector or equivariance; a targeted $\varepsilon$-sweep causally confirms the mechanism ($-2.6$\,pp at $\varepsilon{=}0.1$, $p{<}0.001$). Two five-seed grokking studies confirm the constrained arm does not grok better than the baseline ($p{=}0.019$, A2 wins): the weight-decay exemption has no grokking consequence. A single-seed pilot exploiting this localization achieves the first stable grokking under slingshot conditions---Stiefel + targeted circuit regularization keeps routing-frame isometry error $10^6\times$ lower than the unconstrained ablation through every collapse.
| Comments: | 26 pages, 2 figures |
| Subjects: | Machine Learning (cs.LG); Numerical Analysis (math.NA) |
| MSC classes: | 68T07, 65K10, 53C20, 90C26, 22C05 |
| Cite as: | arXiv:2609.19363 [cs.LG] |
| (or arXiv:2609.19363v3 [cs.LG] for this version) | |
| https://doi.org/10.48550/arXiv.2609.19363 arXiv-issued DOI via DataCite |
Submission history
From: Rubén Darío Guerrero Mr. [view email]
[v1]
Wed, 16 Sep 2026 19:35:47 UTC (147 KB)
[v2]
Sun, 27 Sep 2026 03:19:02 UTC (83 KB)
[v3]
Thu, 1 Oct 2026 15:53:27 UTC (84 KB)
来源:arXiv:cs.LG(机器学习,全量分类) · arxiv.org