拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

STAMP 模型解析:短期注意力与记忆优先机制在会话推荐中的落地实践

STAMP 模型解析:短期注意力与记忆优先机制在会话推荐中的落地实践

1. 从一次线上推荐效果回退说起:STAMP 到底解决什么问题

如果你做过电商或内容平台的推荐系统,大概率遇到过这种场景:用户刚点进来时推荐还算准,但点了三四个商品之后,推荐结果开始"跑偏",越推越像用户很久以前的兴趣,而不是他此刻正在逛的东西。这个问题在 Session-based Recommendation(基于会话的推荐)里特别典型,因为会话本身很短,用户没有历史画像,模型只能靠这一次点击序列来猜他下一步想要什么。

STAMP(Short-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 torch==2.1.0 numpy==1.24.3 pandas==2.0.3 scikit-learn==1.3.0 tqdm==4.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.py:

import torch import torch.nn as nn import torch.nn.functional as F class STAMP(nn.Module): def __init__(self, num_items, embed_dim=100, hidden_dim=100): super(STAMP, self).__init__() self.embedding = nn.Embedding(num_items + 1, embed_dim, padding_idx=0) 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, keepdim=True) # 整体兴趣 # 注意力:每个历史 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 最常用的指标是 Recall@20 和 MRR@20,这两个指标在论文里也是主要对比项。

评估流程是这样的:对测试集里的每个会话,取前 n-1 个 item 作为输入,预测第 n 个 item;模型输出所有候选 item 的得分,取 top-20,看真实 item 是否在里面。代码大致如下:

def evaluate(model, test_loader, topk=20): 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 | Recall@20: 0.512 | MRR@20: 0.221

这个数字在 Diginetica 上属于正常范围。如果你跑出来 Recall@20 只有 0.1 左右,大概率是数据预处理出了问题,比如 item id 映射错位或者 padding 没处理好。

验证阶段还有一个实用技巧:把 STAMP 和 STMP 的评估结果放在一起对比。如果 STAMP 的 Recall@20 比 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_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite

API Key 管理:https://taotoken.net/api-keys?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite

接入文档:https://taotoken.net/doc?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite

如果你打算长期做编码和 Agent 相关的实验,Coding Plan 入口在这里:https://taotoken.net/coding-plan?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite

控制台地址:https://taotoken.net/console?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite

官网首页:https://taotoken.net/?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite

回到模型本身,STAMP 之后可以尝试的改进方向有几个:一是把注意力机制换成多头注意力,看能不能捕捉更细的短期兴趣;二是把 item 的 side information(比如类别、价格)拼进 embedding,缓解冷启动;三是把 STAMP 和 GRU4Rec 做 ensemble,取长补短。这些实验都可以在现有代码基础上改,不需要重写整个框架。

最后提醒一句:复现论文模型时,别急着追 SOTA,先把 baseline 跑稳。STAMP 的价值不在于它有多复杂,而在于它用很轻的结构把"短期兴趣优先"这个直觉落到了实处。你把 STMP 和 STAMP 的对照实验做扎实,比盲目堆模块有用得多。

返回列表