OTKE:基于最优运输的高效序列嵌入方法,对比Transformer自注意力
1. 序列嵌入与特征聚合的核心挑战
在自然语言处理、生物信息学乃至计算机视觉中,我们常常需要处理序列数据。无论是蛋白质的氨基酸序列、句子的单词序列,还是图像经过卷积网络后得到的特征图序列,一个根本性的任务是如何将这些长度不一的序列,转化成一个固定维度的、有意义的向量表示。这个过程,我们称之为序列嵌入或特征聚合。
传统的池化方法,比如均值池化(Mean Pooling)或最大池化(Max Pooling),简单地将所有位置的特征进行平均或取最大值。这种方法计算效率极高,但问题也很明显:它粗暴地抹平了序列内部的顺序和结构信息。想象一下,把一篇影评的所有词向量取平均,得到的向量很可能无法区分“这部电影不算太差”和“这部电影不算太好”的微妙情感差异。
Transformer架构的出现,尤其是其自注意力机制,为这个问题提供了一个优雅的解决方案。自注意力允许序列中的每个元素(例如,一个单词)根据其与所有其他元素的相似度,动态地聚合全局信息。这就像在阅读句子时,大脑会根据当前关注的词,动态地赋予上下文中的其他词不同的重要性。这种机制极大地提升了模型对长程依赖和复杂结构的建模能力,成为如今大语言模型的基石。
然而,自注意力并非没有代价。其计算复杂度是序列长度的平方级(O(n²)),对于长序列(如长文档、基因组序列)来说,内存和计算开销巨大。更重要的是,Transformer的强大表达能力很大程度上依赖于海量的标注数据进行有监督训练。在标注数据稀缺的领域,比如许多生物信息学任务,直接应用Transformer可能面临严重的过拟合风险。
这就引出了一个核心问题: 我们能否设计一种机制,它既具备自注意力那种根据内容动态聚合信息的能力,又能更高效、更适用于数据受限的场景?
最近,一种基于最优运输理论(Optimal Transport, OT)的序列嵌入方法——OTKE,为我们提供了新的思路。它不再让序列进行“自我对齐”,而是引入了一组可学习的“参考点”,让序列与这组参考进行对齐。这种方法在思路上与自注意力有异曲同工之妙,但在计算效率和样本需求上展现出了独特的优势。接下来,我将深入拆解OTKE的工作原理,并将其与Transformer的自注意力进行详细对比,分享我在复现和思考过程中的一些心得。
2. OTKE核心原理:从最优运输到序列嵌入
要理解OTKE,我们得先回到最优运输这个古老的数学问题上。简单来说,最优运输研究的是如何以最小的“成本”,将一堆土(源分布)搬运到另一堆土(目标分布)的位置上。这里的“成本”通常定义为两点间距离的某个函数。
2.1 最优运输的基本思想
在序列嵌入的语境下,我们可以把一个序列看作是一堆“土”。序列中的每个元素(例如,一个词向量或一个图像块的特征)是土堆中的一个颗粒,其质量可以视为均匀或根据重要性分配。OTKE的核心创新在于,它不直接计算两个序列之间的最优运输距离(那依然很昂贵),而是为整个数据集定义了一组共享的、可学习的“参考分布”(可以想象成几个标准的土堆模板)。
对于一个序列x(长度为n)和一个参考分布z(由p个支持点构成,通常p远小于n),我们可以计算从x到z的最优运输计划P(x, z)。这个运输计划是一个n×p的矩阵,其中每一行之和为1/n(假设序列元素质量均匀),每一列之和为1/p。矩阵中的元素P_ij代表了将序列中第i个元素的质量,分配到参考分布第j个支持点上的比例。
为什么这么做是有效的? 这个运输计划P(x, z)本质上是一个软对齐矩阵。它揭示了序列x的结构是如何“映射”到参考模板z上的。如果序列x的某个局部模式与参考z的某个支持点很相似,那么运输计划就会将更多的“质量”分配过去。最终,通过对齐后的特征进行加权求和,我们就得到了序列x关于参考z的固定维度表示Φ_z(x)。
2.2 OTKE与Transformer自注意力的形式化对比
为了更清晰地看到联系与区别,我们可以将OTKE中的关键组件与Transformer的单头自注意力进行并置比较。下面的表格概括了核心差异:
| 特性维度 | Transformer 自注意力 | OTKE (Φ_z) |
|---|---|---|
| 注意力分数计算 | W = Softmax(QK^T/√d), Q,K为查询/键矩阵 | P(x, z), 通过Sinkhorn算法求解最优运输计划 |
| 计算复杂度 | O(n²d) | O(npd), 其中 p ≪ n |
| 对齐对象 | 序列自身内部元素(x 与 x) | 序列与外部参考集 z |
| 参数学习 | 查询矩阵Q、键矩阵K需通过监督学习 | 参考点z可通过无监督(如K-means)或有监督学习 |
| 非线性映射 | 前馈神经网络 | 核函数φ(如高斯核)及其Nyström近似ψ |
| 位置编码 | 绝对或相对位置编码,需额外参数(~qd²) | 通过高斯核融入相对位置成本,参数极少 |
| 参数量级 | ~ qd² (q为头数,d为维度) | ~ qpd (p为支持点数) |
| 是否需要监督 | 通常需要 | 可以不需要 |
从上表可以直观看出OTKE在效率上的优势。自注意力需要计算一个n×n的矩阵,而OTKE只需要计算一个n×p的矩阵。在生物序列分析中,n(序列长度)可能达到数千,而p(参考支持点数量)可能只需要几十到几百,这带来了数量级上的计算和内存节省。
注意 :这里的复杂度O(npd)是粗略估算,主要开销在于计算序列元素与所有参考支持点之间的核函数相似度(n×p次),以及运行Sinkhorn迭代求解运输计划。虽然Sinkhorn迭代有额外成本,但由于p很小,整体仍远小于n²。
2.3 共享参考 vs. 自注意力:哲学差异
这引出了两者最根本的哲学差异:
- Transformer(自注意力) :是“自指”式的。每个序列独立计算其内部元素间的关联。模型通过监督信号学习参数矩阵Q和K,从而学会对于当前任务,什么样的内部关联是重要的。它的归纳偏置是“好的表示来自于序列内部元素的动态组合”。
- OTKE :是“对照”式的。所有序列共享同一组(或多组)参考z。序列表示是通过与这组全局参考进行“比对”或“对齐”而产生的。它的归纳偏置是“好的表示可以通过与一组有代表性的模板进行比较而获得”。
这种“共享参考”的设计带来了几个直接好处:
- 参数效率 :参考点z的维度是p×d,而自注意力中Q/K的维度是n×d(考虑到可变长度,实际参数是d×d)。当p远小于典型的序列长度n时,OTKE参数更少。
- 无监督学习潜力 :参考点z可以通过无监督方式学习,例如对大量未标注序列特征进行聚类(K-means)或计算Wasserstein重心。这意味着我们可以在没有标签的情况下,为特定领域(如蛋白质序列)学习到一组有意义的“模式模板”。
- 可解释性 :学习到的参考点z可以直观地解释为数据集中反复出现的“原型模式”。在生物信息学中,这可能对应着特定的蛋白质结构模体;在NLP中,可能对应着某种句法或语义功能单元。
一个有趣的视角 :实际上,OTKE也可以模拟出自注意力式的行为。如果我们为序列x中的每个元素x_i都单独计算一个到参考z的运输计划P(x_i, z),那么P(x_i, z)P(x_j, z)^T这个矩阵,就可以被解释为元素x_i和x_j之间的一种基于共同参考的相似度度量,这构成了一个低秩的自注意力矩阵近似。这也将OTKE与一系列旨在降低自注意力复杂度的“高效Transformer”工作联系了起来。
3. 实操要点:如何实现一个OTKE层
理解了原理,我们来看看如何动手实现它。这里我会结合论文中的细节和我自己的理解,梳理出关键步骤和注意事项。
3.1 输入与输出定义
假设我们有一个批次(Batch)的序列数据,形状为 (batch_size, n, d) ,其中n是序列长度(可变,需padding),d是特征维度。我们的目标是输出一个固定维度的表示,形状为 (batch_size, output_dim) 。在OTKE中, output_dim = q * p * d' ,其中q是参考的数量,p是每个参考中支持点的数量,d'是经过非线性映射φ后的特征维度(如果φ是恒等映射,则d'=d)。
3.2 关键步骤拆解
3.2.1 参考点(z)的初始化与学习
参考点z是OTKE层的可学习参数,形状为 (q, p, d) 。初始化至关重要。
- 无监督设置 :在训练开始前,我们可以用所有训练数据(或一个子集)的特征,通过K-means聚类算法来初始化z。每个聚类中心就构成了一个参考点的一个支持点。论文中也尝试了用Wasserstein重心来更新,但发现K-means效果相当且更快。
- 有监督设置 :可以将z初始化为随机小数值,或使用无监督方法得到的z作为预训练起点,然后与整个模型一起进行端到端的梯度下降优化。
实操心得 :对于生物序列等数据,使用无监督K-means初始化能提供一个非常好的起点,加速模型收敛,有时甚至能直接得到有竞争力的结果。随机初始化虽然可行,但可能需要更仔细的学习率调整和更长的训练时间。
3.2.2 成本矩阵(Cost Matrix)计算
这是OTKE计算的核心。我们需要计算序列x中每个元素与参考z中每个支持点之间的“距离”或“不相似度”。这个距离由两部分构成:
- 特征成本 :使用一个核函数κ的负值,例如高斯核:
C_feat[i,k] = -exp(-||x_i - z_k||^2 / (2σ^2))。这里σ是带宽参数,控制着相似度衰减的速度。 - 位置成本 :为了捕捉序列顺序,OTKE巧妙地引入了位置编码。它不像Transformer那样将位置信息加到特征上,而是将其融入成本中。例如,使用高斯位置编码:
C_pos[i,k] = exp(-(i/n - α_k)^2 / (2σ_pos^2)),其中α_k是参考点z_k的“虚拟位置”(可学习或固定),σ_pos控制位置平滑度。最终的成本矩阵是特征成本与位置成本的加权和或乘积的负值:C[i,k] = - (C_feat[i,k] * C_pos[i,k])。
为什么这样设计位置编码? 这相当于在最优运输问题中定义了一个新的“地面成本”:两个点之间的运输成本,不仅取决于它们特征上的差异,还取决于它们在序列中位置的差异。这比简单的绝对位置加法更具可解释性,并且参数极少(只有σ_pos和可能的α_k)。
3.2.3 求解最优运输计划(P)
给定成本矩阵C(形状 (n, p) ),我们需要求解最优运输计划P。精确求解OT问题计算量很大,通常采用熵正则化的Sinkhorn算法进行快速近似。该算法通过迭代行、列归一化来求解。
# 简化版的Sinkhorn算法伪代码
def sinkhorn(cost_matrix, epsilon, num_iters):
# cost_matrix: [n, p]
# epsilon: 熵正则化系数
# 初始化
K = torch.exp(-cost_matrix / epsilon) # Gibbs核
u = torch.ones(n, 1) / n
v = torch.ones(1, p) / p
for _ in range(num_iters):
# 行归一化 (对应更新u)
u = 1.0 / (K @ v.T) / n
# 列归一化 (对应更新v)
v = 1.0 / (u.T @ K) / p
# 计算最终的运输计划 P = diag(u) * K * diag(v)
P = torch.diag(u.squeeze()) @ K @ torch.diag(v.squeeze())
return P
- 熵正则化系数ε :这是一个关键超参数。ε越大,得到的运输计划P越平滑(接近均匀分布),计算越稳定但近似误差越大;ε越小,越接近精确OT解,但算法可能不稳定,需要更多迭代。论文中发现在不同任务上,ε=0.5是一个比较稳定有效的值。
- 迭代次数 :论文发现,在监督训练中,即使很少的迭代次数(如10次)也能取得不错的效果,这能显著加速训练。
3.2.4 特征聚合与非线性映射
得到运输计划P后,我们就可以进行聚合了。对于每个参考z_j,计算: aggregated_feature_j = sum_over_i( P_ij * φ(x_i) ) 其中φ是一个非线性映射函数。如果φ是恒等映射,那就是简单的线性加权和。更常见的是使用高斯核等非线性核,这时需要用到Nyström方法来近似这个无限维的映射。
Nyström近似 :为了近似核函数φ,我们随机选取或通过聚类选取一组“锚点”(anchors)w。然后,对于输入x_i,我们可以计算一个有限维的近似特征映射ψ(x_i)。这样, φ(x_i) 就被近似为 ψ(x_i) ,使得整个计算可行。锚点的数量是一个超参数,论文中发现,在无监督设置中,锚点越多近似越好,性能会饱和;在有监督设置中,较小的锚点数就足够了,类似于神经网络中隐藏层神经元的概念。
最后,将所有q个参考聚合后的特征(每个是p×d‘维)拼接起来,再经过一个可选的线性层或直接使用,就得到了序列的最终表示Φ_z(x)。
3.3 参数选择经验
根据论文中的实验,以下是一些参数选择的经验性指导:
- 参考数量q :在无监督生物序列任务中,单个参考(q=1)往往就足够了。在有监督设置中,使用多个参考(如q=5)可能带来轻微提升,但q=1仍然是一个强基线。这类似于Transformer中的多头注意力,但OTKE对头数不那么敏感。
- 支持点数量p :性能随着p增加而提升,但会逐渐饱和。需要权衡表达能力和计算开销。在蛋白质折叠任务中,p=100左右似乎是个甜点。
- 熵正则化ε :在Sinkhorn算法中,ε=0.5在多个任务中表现稳定。
- 位置编码带宽σ_pos :需要根据任务调整。在图像任务(CIFAR-10)中,σ_pos在0.7-1.0之间效果较好;在基因组序列任务(DeepSEA)中,更小的值(如0.1)更有效,这可能因为基因组序列中位置信息非常精确。
4. 实验解析与性能对比
论文在多个领域验证了OTKE的有效性,我们重点看几个有代表性的实验,并分析其背后的原因。
4.1 蛋白质折叠分类(SCOP 1.75)
这是一个经典的生物信息学任务,数据量相对较少(约1.6万训练样本),类别多(1195种折叠),序列长度变化大。这正是OTKE发挥优势的场景。
- 无监督设置 :OTKE在无监督情况下(使用K-means初始化参考点),仅用一个线性分类器,就取得了85.8%的Top-1准确率。这甚至 超过了部分有监督的基线模型 (如Set Transformer的79.2%,ApproxRepSet的84.5%)。这强烈证明了OTKE在无监督或小样本学习中的潜力。其性能优于传统的卷积核网络(CKN,81.8%),主要归功于自适应池化比全局平均池化保留了更多信息。
- 有监督设置 :当进行端到端有监督训练后,OTKE达到了88.7%的Top-1准确率,超越了所有对比的基线模型,包括更深的CNN(DeepSF,73.0%)和循环核网络(RKN,85.3%)。值得注意的是,OTKE的参数量远小于这些模型。
为什么OTKE在这里表现这么好?
- 数据效率 :蛋白质折叠模式往往由一些保守的局部结构模体(motif)组合而成。OTKE学习的参考点,很可能就对应着这些关键的模体。通过与这些模体进行最优运输对齐,模型能够捕捉到序列中哪些区域包含了重要的折叠信息,并进行有针对性的聚合。
- 对长度变化的鲁棒性 :最优运输计划对序列长度不敏感,它关注的是特征的“分布”与参考的匹配程度,而不是绝对位置。这非常适合长度差异巨大的蛋白质序列。
- 计算优势 :对于长序列,自注意力的O(n²)计算难以承受,而OTKE的O(np)复杂度使其能够处理更长的序列。
4.2 情感分析(SST-2)
这是一个标准的NLP任务,数据量相对较大(约7万条评论),序列较短。在这个Transformer占绝对统治地位的领域,OTKE的表现颇具启发性。
OTKE在BERT预训练特征的基础上,仅添加一层嵌入层和一个线性分类器,在无监督和有监督设置下分别达到了86.8%和88.1%的准确率。虽然略低于完全微调的BERT(90.3%),但已经显著超过了简单的 [CLS] 标记(84.6%)或均值池化(85.3%)等基线方法。
这个实验说明了什么? 它表明,即使在Transformer表现优异的领域,OTKE作为一种轻量级的特征聚合/池化层,仍然具有很强的竞争力。特别是考虑到OTKE无需像Transformer那样进行昂贵的预训练或精细的微调。对于资源受限或需要快速原型开发的应用,OTKE是一个非常有吸引力的选择。此外,实验中将OT中的相似度替换为简单的点积后,性能下降明显(从88.1%降至86.9%),这证实了最优运输对齐比简单的线性投影更能捕捉复杂的结构关系。
4.3 计算效率对比
论文中虽然没有给出详细的FLOPs对比,但从算法复杂度和运行时间上可以窥见一斑。在SCOP蛋白质折叠分类的监督实验中:
- Set Transformer :运行时间约3.3小时。
- ApproxRepSet :运行时间约2小时。
- OTKE (5×10) :运行时间约4小时。
- CKN :运行时间约1.5小时。
OTKE比CKN慢,但比Set Transformer快,与ApproxRepSet处于同一量级。考虑到OTKE取得了更好的准确率,这个时间开销是可以接受的。更重要的是,OTKE的内存占用优势是巨大的,因为它不需要存储n×n的注意力矩阵,只需要n×p的运输计划矩阵。
5. 常见问题与调优策略
在实际尝试复现或应用OTKE时,你可能会遇到以下问题,以下是一些排查思路和技巧:
5.1 训练不稳定或梯度爆炸/消失
- 可能原因1:熵正则化系数ε太小 。ε控制着Sinkhorn算法的平滑度。ε过小会使运输计划P变得非常稀疏且尖锐,导致梯度不稳定。
- 解决 :从较大的ε开始(如1.0或0.5),逐步调小观察效果。论文中0.5是一个稳健的起点。
- 可能原因2:成本矩阵C的值范围过大 。在计算Gibbs核
K = exp(-C/ε)时,如果C/ε的值非常大,exp运算可能导致数值溢出(Inf)或下溢(0),使得Sinkhorn迭代失败。- 解决 :对输入特征x和参考点z进行标准化(如LayerNorm或BatchNorm)。确保成本矩阵C的值在一个合理的范围内。可以在计算exp前对C进行裁剪或缩放。
- 可能原因3:Sinkhorn迭代次数不足 。在训练初期,参数变化大,可能需要更多迭代来收敛到一个合理的运输计划。
- 解决 :在训练初期使用较多的迭代次数(如50-100),随着训练稳定,可以逐步减少以加速。论文提到在监督训练中,10次迭代就足够了,但这可能依赖于良好的初始化。
5.2 模型性能不如预期
- 可能原因1:参考点z初始化不佳 。特别是在无监督或小数据场景下,z的初始化至关重要。
- 解决 :务必使用无监督聚类(如K-means)对训练数据特征进行初始化。随机初始化在复杂任务上很难收敛到好结果。
- 可能原因2:位置编码带宽σ_pos设置不当 。σ_pos决定了位置信息的“模糊”程度。对于位置非常关键的任务(如DNA序列),需要小的σ_pos;对于位置相对灵活的任务(如图像块),需要更大的σ_pos。
- 解决 :在验证集上进行网格搜索。可以尝试
[0.05, 0.1, 0.2, 0.5, 1.0]这样的范围。
- 解决 :在验证集上进行网格搜索。可以尝试
- 可能原因3:核函数带宽σ或非线性映射φ选择不当 。
- 解决 :如果特征维度较高或特征值范围大,高斯核的带宽σ需要仔细调整。可以尝试使用线性核(φ为恒等映射)作为基线。对于Nyström近似,确保使用了足够多的锚点来近似核函数。
- 可能原因4:输出维度(q×p×d‘)过大或过小 。这直接影响了模型的容量。
- 解决 :p和q是控制模型容量的主要旋钮。从较小的值开始(如q=1, p=10/50/100),根据验证集性能逐步增加。性能饱和后,增加p或q的收益会很小。
5.3 如何与现有模型结合
OTKE可以作为一个灵活的池化层或特征增强模块插入现有架构。
- 替代CNN/Transformer中的池化层 :在CNN的卷积层之后,可以用OTKE层替代全局平均池化层。在Transformer中,可以用OTKE层替代自注意力层或作为其补充。
- 与预训练模型结合 :如论文实验所示,可以在BERT等预训练模型提取的特征之上,添加一个OTKE层来获得更好的序列表示,然后接一个简单的分类器。这是一种高效的迁移学习策略,无需微调整个大模型。
- 处理图数据 :论文的展望部分提到,将图节点视为序列,OTKE可以作为一种全局图池化方法,聚合图中所有节点的信息,这与基于邻域聚合的经典图神经网络(GNN)形成了不同的归纳偏置。
OTKE为我们提供了一种基于最优运输的、高效且可解释的序列嵌入新工具。它尤其适合那些数据标注成本高、序列长度长、或需要模型具备一定可解释性的领域,如计算生物学、基因组学、医疗时间序列分析等。虽然它可能不会完全取代Transformer在通用序列建模上的地位,但它无疑为我们提供了一个更轻量、更数据高效、且原理优美的替代选择,特别是在资源受限或领域特定的应用中。
更多推荐


所有评论(0)