ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

InterFormer:面向跨序列交互的Transformer变体架构解析

InterFormer:面向跨序列交互的Transformer变体架构解析 1. 从“InterFormer”这个名字说起它到底想解决什么问题第一次看到“InterFormer”这个词很多人会下意识地把它拆成“Inter”和“Former”两半。这个直觉是对的而且恰好点中了它的核心设计哲学。“Former”显然指向 Transformer 架构——过去几年里从自然语言处理到计算机视觉再到多模态融合几乎所有重要突破都建立在这个基础之上。而“Inter”这个前缀通常意味着“交互”“内部”“之间”这几层含义。把两者拼在一起你就能大致猜到它的定位一个专门处理“交互关系”的 Transformer 变体。但这里有个关键问题需要先想清楚为什么现有的 Transformer 不够用非要再搞一个“InterFormer”出来要回答这个问题得回到 Transformer 最核心的机制——自注意力Self-Attention。标准自注意力做的事情是让序列中每个位置去“看”其他所有位置然后根据相关性加权聚合信息。这个机制在单一序列内部非常有效比如一句话里的词与词之间、一张图里的 patch 与 patch 之间。可一旦场景变成“两个或多个序列之间的交互”比如推荐系统里的用户行为序列和商品特征序列、多模态里的文本序列和图像序列、图结构里的节点序列和边序列标准自注意力的处理方式就显得有些“粗暴”了。常见的做法是把多个序列直接拼接成一条长序列然后让自注意力在这条长序列上统一计算。这样做的问题在于拼接会引入大量无意义的跨序列注意力计算。用户行为序列里的某个点击可能跟商品特征序列里的某个属性毫无关系但模型仍然会分配计算资源去算它们的注意力权重。更麻烦的是当序列长度差异很大时比如用户序列有 1000 个行为商品序列只有 50 个特征拼接后的注意力矩阵会变得极度稀疏且不均衡训练效率和效果都会打折扣。InterFormer 的思路就是把“序列内部”和“序列之间”的交互拆开处理。它不再强行把多个序列揉成一条而是设计了一种分层的交互机制先让每个序列在自己的内部做自注意力提取各自的上下文表示然后再用一个专门的“交互模块”去建模序列与序列之间的关联。这个交互模块才是 InterFormer 真正区别于普通 Transformer 的地方。它通常会引入一个可学习的交互矩阵或者交叉注意力机制让不同序列之间的信息流动变得可控、可解释。从应用场景来看InterFormer 这类架构最典型的落地领域包括推荐系统用户行为序列与物品特征序列的交互建模是 CTR 预估、序列推荐等任务的核心。多模态学习文本、图像、音频等不同模态序列之间的对齐与融合。图神经网络节点序列与边序列的交互可以看作是一种特殊的 InterFormer 结构。生物信息学蛋白质序列与配体序列的交互预测比如分子对接、药物筛选。如果你正在做上述任何一个方向的工作并且发现标准 Transformer 在跨序列建模上“力不从心”那 InterFormer 的设计思路就值得你花时间研究。它不是一个凭空造出来的新名词而是对“如何更高效地建模交互关系”这个老问题的系统性回答。2. InterFormer 的核心架构拆解分层交互到底怎么实现2.1 序列内编码先把各自的事情搞清楚InterFormer 的第一步是让每个输入序列独立地经过一个标准的 Transformer 编码器。这一步看起来平平无奇但它的意义在于为后续的交互提供一个干净的、已经包含上下文信息的表示基础。举个例子假设你有一个用户行为序列[点击A, 收藏B, 购买C]和一个商品特征序列[价格, 品牌, 类别]。如果直接拼接商品特征里的“价格”可能会被用户行为里的“点击A”干扰因为自注意力会无差别地计算它们之间的权重。而先做序列内编码用户行为序列会先自己消化掉“点击A→收藏B→购买C”这个行为链条的语义商品特征序列也会先自己理清“价格、品牌、类别”之间的结构关系。两个序列各自“想明白”之后再拿去做交互信息质量会高很多。这一步的实现细节有几个值得注意的地方。位置编码是必须的因为 Transformer 本身没有顺序概念。对于用户行为序列时间戳信息可以编码进去对于商品特征序列如果特征之间有层级关系比如“品牌”属于“类别”也可以用可学习的位置嵌入来表示。另外序列内编码的层数需要根据任务复杂度来定。如果序列本身很短比如只有几个特征一两层就够了如果序列很长比如用户有上千个行为可能需要 4 到 6 层才能充分提取上下文。层数太多会导致过拟合层数太少则上下文信息提取不充分这个平衡点需要通过实验来调。还有一个容易被忽略的细节不同序列的编码器是否共享参数。如果用户序列和商品序列的语义空间差异很大共享参数可能会让模型难以同时学好两个序列的表示。这时候可以考虑用独立的编码器或者共享底层、分离顶层。如果两个序列的语义空间比较接近比如都是文本序列共享参数则可以减少参数量、加速训练。这个选择没有绝对的对错取决于具体任务的数据分布。2.2 交互模块InterFormer 的灵魂所在序列内编码完成后每个序列都得到了一组上下文表示。接下来就是 InterFormer 最核心的部分——交互模块。这个模块的设计目标很明确让不同序列之间的信息能够有选择地流动而不是像拼接自注意力那样“全连接”式地乱流。一种常见的实现方式是交叉注意力Cross-Attention。具体来说把序列 A 作为 Query序列 B 作为 Key 和 Value计算 A 中每个位置对 B 中所有位置的注意力权重然后加权聚合 B 的信息到 A 上。反过来再做一次让 B 也能吸收 A 的信息。这样一轮下来两个序列就完成了一次双向的信息交换。交叉注意力的好处是计算复杂度可控如果 A 的长度是 mB 的长度是 n那么交叉注意力的计算量是 O(m×n)而拼接自注意力是 O((mn)²)。当 m 和 n 都很大时这个差距非常明显。但交叉注意力也有它的局限。如果序列数量超过两个比如多模态场景里有文本、图像、音频三个序列两两做交叉注意力的次数会急剧增加。这时候就需要引入更高效的交互机制比如用一个共享的交互令牌Interaction Token来汇聚所有序列的信息或者设计一个低秩的交互矩阵来近似全连接交互。InterFormer 的论文里通常会讨论这些变体实际使用时需要根据序列数量和计算预算来选择。另一个关键设计是交互的层数。一次交叉注意力只能捕捉一阶交互关系如果序列 A 的某个位置需要经过序列 B 的中转才能影响到序列 C那就需要多层交互。但层数太多又会带来过平滑问题——所有序列的表示会趋于一致失去各自的特性。实践中2 到 3 层交互通常就够了再多就需要配合残差连接和层归一化来稳定训练。2.3 输出融合怎么把交互后的表示拼回一个结果交互模块的输出是每个序列经过信息交换后的新表示。接下来的问题是怎么把这些表示融合成一个最终的预测结果。常见的方法有几种池化融合对每个序列的表示做平均池化或最大池化得到固定长度的向量然后拼接起来送进全连接层。这种方法简单直接但会丢失序列内部的细粒度信息。注意力融合用一个可学习的 Query 去对所有序列的所有位置做注意力自动学习哪些位置对最终预测最重要。这种方法更灵活但参数量会多一些。门控融合为每个序列学习一个门控权重动态决定每个序列对最终结果的贡献比例。这种方法在序列重要性差异很大的场景下特别有用。选择哪种融合方式取决于你的任务需要多细粒度的信息。如果是点击率预估这种只需要一个概率值的任务池化融合通常就够了如果是序列标注这种需要每个位置都有输出的任务就需要更精细的融合方式。3. 动手实现一个最小可用的 InterFormer3.1 环境准备与依赖选择要跑通一个 InterFormer 的最小实现你不需要一上来就搞分布式训练或者大规模预训练。一台带单卡 GPU 的机器加上 PyTorch 和几个常用库就足够验证核心思路了。我建议的依赖清单如下pip install torch2.0.0 pip install numpy pip install tqdm pip install tensorboardPyTorch 2.0 之后的版本对 Transformer 相关操作做了不少优化特别是torch.nn.MultiheadAttention和torch.nn.TransformerEncoderLayer的性能提升明显。如果你用的是更早的版本也能跑但训练速度可能会慢一些。Tensorboard 不是必须的但用来观察 loss 曲线和注意力权重分布非常方便建议装上。数据方面如果你手头没有现成的多序列交互数据集可以先用一个简单的合成数据集来验证代码正确性。比如构造两个序列序列 A 是随机生成的 0/1 向量序列序列 B 是随机生成的浮点数序列标签是序列 A 的某个位置和序列 B 的某个位置的组合函数。这样你就能清楚地知道模型有没有学到交互关系。3.2 序列内编码器的代码实现先写序列内编码器。这里我用一个简化的 Transformer Encoder 来实现核心是自注意力加前馈网络import torch import torch.nn as nn import math class SequenceEncoder(nn.Module): def __init__(self, d_model, nhead, num_layers, dim_feedforward, dropout0.1): super().__init__() self.d_model d_model encoder_layer nn.TransformerEncoderLayer( d_modeld_model, nheadnhead, dim_feedforwarddim_feedforward, dropoutdropout, batch_firstTrue ) self.encoder nn.TransformerEncoder(encoder_layer, num_layersnum_layers) self.pos_embedding nn.Embedding(512, d_model) # 假设最大序列长度512 def forward(self, x, maskNone): # x: (batch_size, seq_len, d_model) seq_len x.size(1) positions torch.arange(seq_len, devicex.device).unsqueeze(0) x x self.pos_embedding(positions) return self.encoder(x, src_key_padding_maskmask)这段代码里有两个细节值得展开。第一是位置编码我用的是可学习的位置嵌入而不是正弦余弦编码。可学习嵌入的好处是模型可以根据任务数据自动调整位置表示特别是在序列长度不固定、位置语义比较复杂的场景下比如用户行为序列里最近的行为和很久以前的行为位置关系不是简单的线性距离可学习嵌入通常效果更好。第二是 mask 的处理src_key_padding_mask用来屏蔽 padding 位置避免它们参与注意力计算。如果你的序列长度是固定的可以不用 mask但如果用了 padding这个 mask 必须传进去否则 padding 位置的随机值会污染注意力权重。3.3 交互模块的代码实现交互模块是 InterFormer 的核心我用交叉注意力来实现class InteractionModule(nn.Module): def __init__(self, d_model, nhead, dropout0.1): super().__init__() self.cross_attn_a2b nn.MultiheadAttention(d_model, nhead, dropoutdropout, batch_firstTrue) self.cross_attn_b2a nn.MultiheadAttention(d_model, nhead, dropoutdropout, batch_firstTrue) self.norm_a nn.LayerNorm(d_model) self.norm_b nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, seq_a, seq_b, mask_aNone, mask_bNone): # seq_a: (batch, len_a, d_model), seq_b: (batch, len_b, d_model) # A 吸收 B 的信息 attn_out_a, _ self.cross_attn_a2b( queryseq_a, keyseq_b, valueseq_b, key_padding_maskmask_b ) seq_a self.norm_a(seq_a self.dropout(attn_out_a)) # B 吸收 A 的信息 attn_out_b, _ self.cross_attn_b2a( queryseq_b, keyseq_a, valueseq_a, key_padding_maskmask_a ) seq_b self.norm_b(seq_b self.dropout(attn_out_b)) return seq_a, seq_b这里有几个实操中容易踩的坑。第一个坑是 mask 的方向。在cross_attn_a2b里Query 是 seq_aKey 和 Value 是 seq_b所以key_padding_mask应该传 seq_b 的 mask而不是 seq_a 的。这个很容易搞反一旦搞反模型会去注意 padding 位置训练 loss 会异常震荡。第二个坑是残差连接和层归一化的顺序。我用的是 Post-Norm先残差再加归一化这是原始 Transformer 的做法。但后来很多工作发现 Pre-Norm先归一化再进注意力训练更稳定特别是在深层网络里。如果你的交互模块堆了很多层建议换成 Pre-Norm。第三个坑是 dropout 的位置。我在残差连接前加了 dropout这是标准做法。但如果你发现模型欠拟合可以把 dropout 调小或者去掉如果过拟合严重可以适当增大。3.4 完整模型组装与训练循环把序列编码器和交互模块拼起来再加上输出层就是一个完整的 InterFormerclass InterFormer(nn.Module): def __init__(self, d_model128, nhead4, num_encoder_layers2, num_interaction_layers2, dim_feedforward256, dropout0.1): super().__init__() self.encoder_a SequenceEncoder(d_model, nhead, num_encoder_layers, dim_feedforward, dropout) self.encoder_b SequenceEncoder(d_model, nhead, num_encoder_layers, dim_feedforward, dropout) self.interaction_layers nn.ModuleList([ InteractionModule(d_model, nhead, dropout) for _ in range(num_interaction_layers) ]) self.output_layer nn.Sequential( nn.Linear(d_model * 2, d_model), nn.ReLU(), nn.Dropout(dropout), nn.Linear(d_model, 1) ) def forward(self, seq_a, seq_b, mask_aNone, mask_bNone): # 序列内编码 enc_a self.encoder_a(seq_a, mask_a) enc_b self.encoder_b(seq_b, mask_b) # 多层交互 for layer in self.interaction_layers: enc_a, enc_b layer(enc_a, enc_b, mask_a, mask_b) # 池化融合 pooled_a enc_a.mean(dim1) # (batch, d_model) pooled_b enc_b.mean(dim1) combined torch.cat([pooled_a, pooled_b], dim-1) return self.output_layer(combined)训练循环用标准的 Adam 优化器和二元交叉熵损失model InterFormer().cuda() optimizer torch.optim.Adam(model.parameters(), lr1e-3, weight_decay1e-5) criterion nn.BCEWithLogitsLoss() for epoch in range(50): model.train() total_loss 0 for batch in train_loader: seq_a, seq_b, labels batch seq_a, seq_b, labels seq_a.cuda(), seq_b.cuda(), labels.cuda() optimizer.zero_grad() logits model(seq_a, seq_b).squeeze(-1) loss criterion(logits, labels) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() total_loss loss.item() print(fEpoch {epoch}, Loss: {total_loss / len(train_loader):.4f})这里我加了梯度裁剪clip_grad_norm_因为 Transformer 类模型在训练初期容易出现梯度爆炸特别是交互模块里的交叉注意力梯度回传路径比较长。裁剪阈值设成 1.0 是个比较安全的起点如果发现训练不稳定可以再调小。4. 训练 InterFormer 时最容易踩的五个坑4.1 序列长度差异过大导致的注意力失衡这是我在实际项目里遇到的第一个大坑。当时用户行为序列平均长度是 200商品特征序列平均长度是 15两者差了十几倍。直接做交叉注意力时模型几乎把所有注意力都放在了用户序列上商品序列的信息被严重稀释。后来我做了两件事来解决第一是在交叉注意力里给 Key 和 Value 加一个长度归一化让注意力权重的总和不受序列长度影响第二是在池化融合时给短序列更高的权重因为短序列的信息密度通常更高。这两个调整之后商品特征对最终预测的贡献明显提升了。具体实现上长度归一化可以通过在 softmax 之前除以sqrt(key_length)来实现类似缩放点积注意力里的缩放因子。池化融合的权重可以用一个可学习的标量初始值设为1/sqrt(seq_len)让模型自己学。4.2 交互层数过多导致的表示坍缩前面提到过交互层数太多会让所有序列的表示趋于一致。我试过堆 6 层交互结果发现两个序列的表示余弦相似度超过了 0.95几乎变成同一个向量了。这时候模型虽然训练 loss 很低但验证集效果很差因为序列的独特信息全丢了。解决方案是引入一个“交互强度”门控每一层交互的输出都跟原始序列表示做一个加权平均权重由模型自己学。这样即使交互层数多模型也能保留一部分原始信息。另一个方案是在交互模块里加一个可学习的温度系数控制注意力分布的尖锐程度温度越高注意力越分散信息融合越温和。4.3 位置编码在交互后的错位问题序列内编码时位置编码是加在原始输入上的。但经过交叉注意力之后序列 A 的每个位置都吸收了序列 B 的信息这时候原来的位置编码还准确吗答案是不再准确了但通常不需要重新加。因为交叉注意力本身不改变序列的长度和顺序位置信息在注意力计算中已经通过 Query 和 Key 的对应关系隐式保留了。如果你发现模型对位置特别敏感比如序列推荐里最近的行为权重应该更高可以在交互后再加一层可学习的位置偏置让模型自己调整。4.4 训练数据中交互信号的稀疏性很多真实场景下两个序列之间的交互信号是非常稀疏的。比如用户行为序列里只有少数几个行为跟商品特征有强关联大部分行为是噪声。这时候如果直接用交叉注意力模型会被大量弱交互信号淹没。我的做法是在交互模块前加一个“交互筛选”层用一个轻量的 MLP 给每个位置对打一个交互分数只保留 top-k 的交互对进入交叉注意力。这个筛选层可以和主模型一起端到端训练筛选阈值用 Gumbel-Softmax 来保证可微。实测下来这个方法在稀疏交互场景下能提升 3 到 5 个点的 AUC。4.5 推理时的计算效率优化InterFormer 在训练时可以用完整的交叉注意力但推理时如果序列很长计算量会成为瓶颈。一个实用的优化是缓存序列内编码的结果。因为序列内编码不依赖另一个序列可以在用户请求到来之前预先算好并缓存。推理时只需要跑交互模块和输出层延迟能降低 60% 以上。另一个优化是对交互模块做量化把交叉注意力的计算从 FP32 降到 INT8精度损失通常在 1% 以内但速度能提升 2 到 3 倍。这两个优化我在线上环境都验证过效果很稳。5. InterFormer 在不同场景下的适配策略5.1 推荐系统用户序列与物品序列的交互推荐系统是 InterFormer 最自然的应用场景。用户行为序列和物品特征序列的交互本质上是在回答“这个用户对这个物品感兴趣的概率有多大”。但这里有个特殊之处用户序列是动态的物品序列是静态的。同一个物品会被很多用户看到但每个用户的行为序列都是独特的。所以实践中物品序列的编码可以离线预计算并缓存用户序列的编码和交互则在线实时计算。这种“静态缓存动态计算”的架构在工业级推荐系统里非常常见。另外推荐系统里的交互往往不是对称的。用户行为对物品表示的影响和物品特征对用户表示的影响重要程度可能完全不同。比如在电商场景里用户的历史购买行为对理解物品的卖点很重要但物品的价格特征对理解用户偏好可能没那么关键。这时候可以用非对称的交互模块给两个方向的交叉注意力不同的容量比如用户→物品用 4 头注意力物品→用户用 2 头让模型把更多参数花在更重要的交互方向上。5.2 多模态学习文本、图像、音频的三方交互多模态场景下序列数量通常超过两个两两交叉注意力的计算量会爆炸。这时候可以用共享交互令牌的方案引入一组可学习的令牌每个模态的序列都先跟这组令牌做交叉注意力令牌再跟其他模态的令牌做交互。这样计算复杂度从 O(n²) 降到 O(n×k)k 是令牌数量通常设成 8 到 16 就够了。这个方案在视频理解、图文匹配等任务上都有不错的表现。另一个多模态特有的问题是模态间的对齐。文本序列和图像序列的长度和语义粒度往往差异很大直接做交叉注意力效果不好。实践中会先用一个对齐模块把两个序列投影到同一个语义空间比如用对比学习预训练一个共享的嵌入空间然后再做交互。这个对齐步骤对最终效果影响很大值得多花时间调。5.3 图结构数据节点序列与边序列的交互图数据可以看作是一种特殊的 InterFormer 场景节点序列和边序列的交互。但图结构有个特点——交互是稀疏的、局部的。一个节点只跟它的邻居节点有边不是跟所有节点都有边。所以直接用全连接的交叉注意力是不合适的需要引入邻接矩阵掩码把没有边的位置对的注意力权重屏蔽掉。这样既符合图的结构先验又能大幅减少计算量。对于大规模图还可以用邻居采样的策略每个节点只采样固定数量的邻居参与交互而不是用全部邻居。这样计算量就固定了不会随节点度数增加而爆炸。采样数量通常设成 10 到 20太少会丢失信息太多则计算效率下降。6. 一些关于 InterFormer 的常见误解与澄清6.1 它不是 Transformer 的替代品很多人看到“InterFormer”这个名字会以为它是来取代 Transformer 的。其实不是。InterFormer 是 Transformer 在跨序列交互场景下的一个特化变体它的序列内编码部分用的还是标准 Transformer。你可以把它理解成“Transformer 交互模块”的组合。如果你的任务只涉及单一序列用标准 Transformer 就够了没必要上 InterFormer。只有当你需要建模多个序列之间的交互关系时InterFormer 的设计才有意义。6.2 它不一定要用交叉注意力交叉注意力只是实现交互模块的一种方式不是唯一方式。交互模块的核心目标是“有选择地让信息在序列间流动”交叉注意力只是达成这个目标的一种手段。其他可行的手段包括用图神经网络来建模序列间的交互、用门控机制来控制信息流动、用低秩分解来近似全连接交互。选择哪种方式取决于你的数据特点、计算预算和任务需求。交叉注意力在序列长度适中、交互比较稠密的场景下表现最好如果交互非常稀疏图神经网络可能更合适。6.3 它不需要从头训练InterFormer 的序列内编码器可以直接用预训练好的 Transformer 权重来初始化比如 BERT、RoBERTa 或者 Vision Transformer。这样能省掉大量预训练成本特别是在数据量不够大的场景下预训练初始化能显著提升效果。交互模块则需要从头训练因为它没有现成的预训练权重可用。但交互模块的参数量通常不大训练成本可以接受。7. 我个人的一些实操体会做 InterFormer 相关项目这两年最大的体会是交互模块的设计比序列内编码重要得多。很多人把大量精力花在调序列编码器的层数、头数、维度上但真正决定效果上限的是交互模块能不能有效地捕捉到序列间的关键关联。我见过太多项目序列编码器调得很精细但交互模块就是简单的拼接加全连接结果效果一直上不去。后来把交互模块换成交叉注意力同样的序列编码器AUC 直接涨了 4 个点。另一个体会是不要迷信论文里的默认配置。InterFormer 类模型的超参数非常依赖具体任务和数据分布。论文里在某个数据集上效果最好的配置换到你的数据上可能完全不行。我的做法是先用一个很小的配置比如 d_model64nhead2交互层数1跑通流程确认数据管道和训练逻辑没问题然后再逐步放大模型、调超参数。这样能避免一上来就搞大模型、结果训练几天发现数据有问题的情况。最后分享一个调参小技巧交互模块的学习率通常要比序列编码器高一些。因为交互模块是从头训练的而序列编码器可能用了预训练权重需要的学习率更低。我一般会给交互模块设 2 到 3 倍于序列编码器的学习率用参数组的方式在优化器里分开设置。这个技巧在多个项目里都帮我省了不少调参时间。
返回列表