提出新型分块可微Sinkhorn注意力,提升长序列建模精度与效率。
Block-Wise Differentiable Sinkhorn Attention: Tail-Refinement Gradients with a Gap-Aware Dustbin Bridge

- 设计停步-固定深度的尾部精修代理模型,实现精确反向传播
- 在TPU上实现每秒8.5个样本,3小时完成4配置Pfam筛查
- 支持高精度梯度计算,适用于长序列建模与硬件优化场景
本文研究在TPU硬件上基于停止基、固定深度尾部精修代理的长上下文平衡熵最优传输(OT)注意力。在停止$T$步的Sinkhorn求解后,通过回滚短尾精修并精确微分该代理。对于报告的$R=2$ TPU路径,反向传播包含四个阶梯状计划因子。证明了单参考瓷砖的精确调度:$R=2$得分余切为单个参考计划瓷砖乘以由向量余切和对偶差构建的显式修正场。该方法实现分块代价$O((T+R)LW)$,输入存储$O(Ld)$,额外高带宽内存使用$O(L)$,固定头维度$d$与带宽$W$。同时将当前\texttt{dustbin\_block}路径形式化为扩展支持上的同一单位目标代理,使伴随调度可提升至实际使用的单活跃尘箱路径;此桥接为代数结构,不声称适用于一般KL非平衡或任意容量间隙模型。提供局部代理偏差界、后验偏差证书及严格正激活块的投影收缩证书。在合成掩码问题中,优化核与相同中心代理的精确自动微分结果误差在$10^{-5}$至$10^{-10}$之间。在TPU v6e-8上,四配置Pfam筛查全程完成,推广的平衡$R=2$运行在三小时预算内维持约8.5例/秒,达到第1437步。保留测试片段的重构误差从5.57降至2.05,稀疏交叉熵从5.53降至5.30,相对于步骤0;交叉熵作为诊断日志而非直接优化目标;目标重心对齐指标未显著改善,确定性对角参考仍在此类指标上表现更优。
原文摘要 · Abstract (English)
We study long-context balanced entropic optimal transport (OT) attention on TPU hardware through a stopped-base, fixed-depth tail-refinement surrogate. After a stopped $T$-step Sinkhorn solve, we unroll a short refinement tail and differentiate that surrogate exactly. For the reported $R=2$ TPU path, the backward pass contains four staircase plan factors. We prove an exact one-reference-tile schedule: the $R=2$ score cotangent is a single reference plan tile times an explicit modifier field built from vector cotangents and dual differences. This yields block-wise cost $O((T+R)LW)$, $O(Ld)$ input storage, and $O(L)$ additional HBM usage for fixed head dimension $d$ and band width $W$ on the balanced fixed-support path. We also formalize the current \texttt{dustbin\_block} path as the same unit-target surrogate on an augmented support, so the adjoint schedule lifts to the single-active-dustbin path used in our TPU runs; this bridge is algebraic and does not claim a general KL-unbalanced or arbitrary-capacity gap model. We provide a local surrogate-bias bound, an a posteriori bias certificate, and a projective contraction certificate for strictly positive active blocks. On synthetic masked problems, the optimized kernel matches exact autodiff of the same centered surrogate to within $10^{-5}$--$10^{-10}$. On TPU v6e-8, a four-configuration Pfam screen completes end-to-end, and a promoted balanced $R=2$ run sustains roughly $8.5$ examples per second through a three-hour budget, reaching step $1437$. Held-out Pfam test shards improve reconstruction from $5.57$ to $2.05$ and sparse CE from $5.53$ to $5.30$ relative to step $0$, with CE logged diagnostically rather than optimized directly; target-barycenter alignment metrics do not materially improve, and a deterministic diagonal reference remains stronger on those metrics.
Thank you to arXiv for use of its open access interoperability. PaperDance 不是 arXiv 官方产品;中文卡片由大模型生成,请以原文为准。