不能直接套用nlp的selfattention,因din需targetattention:仅单向query→keys聚焦,须显式对齐维度、加padding mask、拼接query-key特征后过mlp生成权重,且需针对稀疏/高相似序列设计fallback与归一化策略。

直接用 TargetAttention 层替代传统 SumPooling,就能让模型对每个候选 item 动态聚焦相关历史行为——但权重崩坏、稀疏失效、特征耦合这三类问题,不提前干预,上线后指标必掉。
为什么不能直接套用 NLP 的 SelfAttention?
DIN 的核心是 TargetAttention,不是 SelfAttention。前者只让用户行为序列(keys)向当前候选 item(query)单向“看”,后者会让所有行为互相“看”。在推荐场景中,用户点击球鞋和浏览口红之间本无时序依赖,强行建模内部关系反而引入噪声。
- 输入维度错配:NLP 中
query/key通常同维;DIN 中query是 target item embedding,keys是用户行为 embedding,二者可能来自不同 embedding 表,需显式对齐或变换 - 缺失 padding mask:用户行为序列长度不一,必须用
mask屏蔽 padding 位置的 attention 权重,否则softmax会把无效位置分走概率 - 无 query-key 交互增强:原始
SelfAttention只算内积;DIN 论文中关键改进是拼接[query, key, query - key, query * key]再过 MLP,保留差异与共性信号
TargetAttention 层的 PyTorch 实现要点
别直接写 torch.nn.MultiheadAttention——它默认做 SelfAttention,且不支持 query-key 特征拼接。应手写一个轻量 nn.Module:
- 输入必须含
keys_length(每个样本的真实序列长度),用于生成mask:用torch.arange(seq_len)与keys_length.unsqueeze(1)比较,得到boolmask 矩阵 - 权重计算推荐用 DIN 原论文的 MLP 方式:
torch.cat([query, keys, query - keys, query * keys], dim=-1)→self.mlp→softmax,比单纯点积更鲁棒 - 输出是加权和:
torch.sum(weights.unsqueeze(-1) * values, dim=1),注意values通常等于keys,但留出接口便于后续替换(如接入 GRU 输出) - 初始化
mlp最后一层 bias 为 0,避免初始权重均匀分布;训练初期用tf.print或torch.histc监控weights分布:均值应在 0.3–0.7,方差 > 0.1
稀疏行为与相似 item 导致的权重失效怎么破?
新用户只有 2–3 次点击,或某用户连续刷了 10 条同款手机视频,TargetAttention 容易输出退化权重(全 0.1 或单峰尖刺),这不是 bug,是机制局限。
- 对低频用户:在 embedding 层前加 fallback 逻辑——若
keys_length ,跳过 attention,直接用 <code>MeanPooling+ 小幅度 dropout,避免权重乱跳 - 对高相似序列:在 MLP 输入前加入
LayerNorm,抑制query * key项因 embedding 过于接近导致的数值饱和 - 硬约束可选:对 softmax 输出加
temperature调节(如除以 0.5),但会削弱区分度;更稳妥的是在 loss 中加辅助项:鼓励 top-3 权重之和 ≥ 0.6 - 别忽略
query质量:target item embedding 若未 fine-tune(比如直接用预训练 ID embedding),query表达能力弱,attention 必然失效——建议至少 joint train item tower
TensorFlow 2.x 里 TargetAttention 的部署陷阱
TF SavedModel 导出后线上 infer 时,常出现权重全为 0 或 nan,90% 源于动态 shape 处理不当。
-
keys输入必须声明shape=[None, None, embedding_dim],不能写死seq_len;否则 batch 内不同长度样本会被 pad 到同一长度,mask 逻辑错位 -
tf.nn.softmax在axis=-1上操作时,若未 mask,padding 位置参与归一化——务必用tf.where(mask, scores, -1e9)替换负无穷,避免 FP16 下 inf 溢出 - 自定义
call方法中,禁止用tf.print输出中间 tensor(影响图模式构建);调试用tf.debugging.assert_all_finite检查scores和weights - 线上服务时,batch size 变化会导致
keys第二维动态变化,Keras Model 需设jit_compile=False,否则 XLA 编译失败
真正难的不是写出 attention 公式,而是让权重在千万级用户、千变万化行为序列下始终稳定有效——mask 是否严丝合缝、query 是否足够有判别力、稀疏 case 是否有 fallback,这三点漏掉任一,模型就只是个好看的离线指标提升器。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











