STAMP 模型解析:短期注意力与记忆优先机制在会话推荐中的落地实践
1. 从一次线上推荐效果回退说起STAMP 到底解决什么问题如果你做过电商或内容平台的推荐系统大概率遇到过这种场景用户刚点进来时推荐还算准但点了三四个商品之后推荐结果开始跑偏越推越像用户很久以前的兴趣而不是他此刻正在逛的东西。这个问题在 Session-based Recommendation基于会话的推荐里特别典型因为会话本身很短用户没有历史画像模型只能靠这一次点击序列来猜他下一步想要什么。STAMPShort-Term Attention/Memory Priority Model就是冲着这个痛点来的。它的核心主张很直接用户的兴趣由两部分组成一部分是这次会话里累积出来的整体兴趣general interest另一部分是他最后一次点击所代表的当前兴趣current interest。传统做法用 LSTM 把整个序列编码成一个隐状态理论上能记住长期依赖但作者在论文里指出LSTM 对长会话的建模其实并不够有效——序列一长早期信息被稀释最后那个隐状态未必能准确反映用户现在想要什么。STAMP 的解法是不再只依赖 LSTM 的最终隐状态而是显式地把会话平均表示和最后一次点击表示都拿出来用一个注意力网络去算每个历史 item 对当前兴趣的贡献权重再加权求和。这样既保留了整体兴趣的稳定性又强化了短期兴趣的优先级。适合谁看如果你正在做推荐系统、想复现一个结构不复杂但效果扎实的 baseline或者你已经在用 LSTM/GRU 做序列推荐但效果卡住了这篇的配置和验证步骤可以直接拿去跑。我试过在公开数据集上把 STAMP 和纯 LSTM 版本做对照差距在短会话上尤其明显。下面从模型结构、配置、训练到排障一步步拆开讲。2. 环境与依赖准备TaoToken 接入前的模型侧配置在真正写 STAMP 之前先把运行环境和依赖理清楚。STAMP 本身是一个相对轻量的模型核心就是 embedding 层、注意力层和 MLP 打分层不需要特别重的框架。我一般用 PyTorch 来复现因为注意力权重的调试比较直观。先建一个干净的虚拟环境把依赖固定下来。这里给出一个可复制的 requirements 片段路径按你自己的项目根目录来# requirements.txt torch2.1.0 numpy1.24.3 pandas2.0.3 scikit-learn1.3.0 tqdm4.66.1安装命令python -m venv venv source venv/bin/activate # Windows 用 venv\Scripts\activate pip install -r requirements.txt数据集方面Session-based Recommendation 最常用的两个公开数据集是 Diginetica 和 Yoochoose现在多叫 RetailRocket 的变体。它们都是会话-点击序列-下一个点击的格式。预处理要做的事很固定按时间戳切会话、过滤掉长度小于 2 的会话、把 item id 重映射成从 1 开始的连续整数0 留给 padding。这里有个容易踩的坑item 重映射一定要在切完训练/测试集之后统一做否则训练集和测试集的 id 空间对不上模型跑起来 loss 会莫名其妙地不降。我一般写一个build_vocab函数先扫全量数据建映射表再分别处理各集合。如果你在团队里做协作模型代码和配置建议放到一个统一的地方管理。我平时会把实验配置、模型权重路径、日志目录都写进一个config.yaml这样换数据集时只改配置不改代码。至于模型训练本身STAMP 对显存要求不高单卡 8G 就能跑中等规模数据集batch size 设 128 或 256 都行。环境准备好之后下一步就是真正把 STAMP 的结构写出来。这里要特别注意STAMP 有两个版本一个是 STMP不带注意力一个是 STAMP带注意力。很多人复现时直接上 STAMP结果发现和论文对不上其实是因为没先跑通 STMP 做对照。建议两个都实现方便验证注意力层到底带来了多少提升。3. 可复制的 STAMP 模型结构与训练配置这一节是核心直接给可运行的模型定义和训练参数。先看 STAMP 的结构逻辑输入是一个会话的 item 序列经过 embedding 层得到每个 item 的向量然后分两路一路对序列做平均得到整体兴趣表示一路取最后一个 item 的向量作为当前兴趣表示接着用注意力机制计算每个历史 item 对当前兴趣的权重加权求和得到短期兴趣表示最后把整体兴趣、当前兴趣、短期兴趣拼接或相加后送入 MLP输出每个候选 item 的得分。下面是一个精简但完整的 PyTorch 实现你可以直接复制到model.pyimport torch import torch.nn as nn import torch.nn.functional as F class STAMP(nn.Module): def __init__(self, num_items, embed_dim100, hidden_dim100): super(STAMP, self).__init__() self.embedding nn.Embedding(num_items 1, embed_dim, padding_idx0) self.attn_mlp nn.Sequential( nn.Linear(embed_dim * 2, hidden_dim), nn.Sigmoid() ) self.fc1 nn.Linear(embed_dim * 3, hidden_dim) self.fc2 nn.Linear(hidden_dim, embed_dim) def forward(self, seq, mask): # seq: [batch, seq_len], mask: [batch, seq_len] emb self.embedding(seq) # [B, L, D] last emb[:, -1, :] # 当前兴趣 [B, D] avg (emb * mask.unsqueeze(-1)).sum(1) / mask.sum(1, keepdimTrue) # 整体兴趣 # 注意力每个历史 item 与当前兴趣的交互 last_exp last.unsqueeze(1).expand_as(emb) # [B, L, D] attn_input torch.cat([emb, last_exp], dim-1) attn_score self.attn_mlp(attn_input).sum(-1) # [B, L] attn_score attn_score.masked_fill(mask 0, -1e9) attn_weight F.softmax(attn_score, dim-1) # [B, L] short (emb * attn_weight.unsqueeze(-1)).sum(1) # 短期兴趣 [B, D] concat torch.cat([avg, last, short], dim-1) out self.fc2(F.relu(self.fc1(concat))) return out, attn_weight对应的训练配置我一般写成一个config.yaml路径和参数都固定下来方便复现data: train_path: ./data/train.txt test_path: ./data/test.txt max_seq_len: 50 model: embed_dim: 100 hidden_dim: 100 train: batch_size: 256 lr: 0.001 epochs: 30 optimizer: Adam loss: CrossEntropyLoss weight_decay: 0.00001训练循环里有个细节要注意STAMP 的损失函数是标准的交叉熵但负样本的构造方式会影响效果。论文里用的是对每个正样本随机采样若干负样本的方式我实测下来如果直接用全量 item 做 softmax计算量大且收敛慢建议先用负采样跑通再考虑全量。另外注意力权重的可视化对调试很有帮助。你可以在验证阶段把attn_weight存下来看看模型是不是真的把高权重给了最近几个点击。如果权重分布很均匀说明注意力层没学到东西可能是学习率太大或者 embedding 维度太小。4. 验证请求与成功结果离线评估怎么跑模型训练完之后必须做离线评估否则你不知道它到底有没有比 baseline 好。Session-based Recommendation 最常用的指标是 Recall20 和 MRR20这两个指标在论文里也是主要对比项。评估流程是这样的对测试集里的每个会话取前 n-1 个 item 作为输入预测第 n 个 item模型输出所有候选 item 的得分取 top-20看真实 item 是否在里面。代码大致如下def evaluate(model, test_loader, topk20): model.eval() recall, mrr, total 0.0, 0.0, 0 with torch.no_grad(): for seq, mask, target in test_loader: scores, _ model(seq, mask) _, topk_idx torch.topk(scores, topk, dim-1) for i in range(target.size(0)): total 1 rank (topk_idx[i] target[i]).nonzero() if rank.numel() 0: recall 1 mrr 1.0 / (rank.item() 1) return recall / total, mrr / total跑通之后你会看到类似这样的输出Epoch 30 | Loss: 2.134 | Recall20: 0.512 | MRR20: 0.221这个数字在 Diginetica 上属于正常范围。如果你跑出来 Recall20 只有 0.1 左右大概率是数据预处理出了问题比如 item id 映射错位或者 padding 没处理好。验证阶段还有一个实用技巧把 STAMP 和 STMP 的评估结果放在一起对比。如果 STAMP 的 Recall20 比 STMP 高 3-5 个点说明注意力层确实起作用了如果两者差不多那就要检查注意力权重是不是退化了。另外评估时要注意测试集的会话长度分布。如果大部分会话都很短比如只有 2-3 个 item那 STAMP 的优势可能不明显因为短期兴趣和整体兴趣几乎重合。这种情况下可以单独统计长会话长度大于 10上的指标更能看出模型差异。5. 本篇常见错排查从 401 到注意力权重异常复现 STAMP 的过程中报错主要集中在几个地方。我把自己踩过的坑列出来对照着排查会快很多。第一个常见错误是RuntimeError: expected scalar type Long but found Float。这通常是因为 embedding 层的输入要求是整数类型的 item id但你在预处理时把 id 转成了 float。解决办法是检查seq的数据类型确保它是torch.long。在 DataLoader 里加一句seq seq.long()就能解决。第二个是IndexError: index out of range in self。这是 embedding 的经典问题你的 item id 最大值超过了num_items。比如你建 vocab 时统计的是训练集但测试集里出现了训练集没有的 item。解决办法是在预处理阶段统一建 vocab或者给未知 item 留一个专门的 id。第三个是注意力权重全为 0 或者全相等。这通常发生在 mask 处理不当的时候。如果你的mask是 bool 类型masked_fill要用mask 0如果是 float 类型要确保 padding 位置确实是 0。我一般会在 forward 里打印一次attn_weight的均值和方差确认它不是一个常数。第四个是 loss 不下降或者震荡。除了学习率太大之外还有一个容易被忽略的原因负采样数量太少。如果每个正样本只采 1 个负样本梯度噪声会很大。建议至少采 5-10 个或者直接用全量 softmax 跑小数据集验证。第五个是评估指标异常低。除了数据预处理问题还要检查评估时是不是把 padding 也当成了候选 item。正确的做法是在计算 top-k 时把 padding 位置的得分设为负无穷。如果你在接入外部服务做实验管理时遇到401 Unauthorized或local proxy failed这类报错通常是鉴权信息没配对。这时候可以检查一下 API Key 是否写进了环境变量以及 Base URL 是否指向了正确的地址。模型侧和平台侧的配置要分开排查别混在一起调。6. 语义一致的接入与后续实验建议把 STAMP 跑通之后下一步通常是把它接入到实际的实验流程里。如果你需要统一管理模型对话、API Key 和编码计划可以按下面的路径操作模型对话调试入口https://taotoken.net/api?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteAPI Key 管理https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite接入文档https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite如果你打算长期做编码和 Agent 相关的实验Coding Plan 入口在这里https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite控制台地址https://taotoken.net/console?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite官网首页https://taotoken.net/?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite回到模型本身STAMP 之后可以尝试的改进方向有几个一是把注意力机制换成多头注意力看能不能捕捉更细的短期兴趣二是把 item 的 side information比如类别、价格拼进 embedding缓解冷启动三是把 STAMP 和 GRU4Rec 做 ensemble取长补短。这些实验都可以在现有代码基础上改不需要重写整个框架。最后提醒一句复现论文模型时别急着追 SOTA先把 baseline 跑稳。STAMP 的价值不在于它有多复杂而在于它用很轻的结构把短期兴趣优先这个直觉落到了实处。你把 STMP 和 STAMP 的对照实验做扎实比盲目堆模块有用得多。