先機(jī)制在會(huì)話(huà)推薦中的落地實(shí)踐)
1. 從一次線(xiàn)上推薦效果回退說(shuō)起STAMP 到底解決什么問(wèn)題如果你做過(guò)電商或內(nèi)容平臺(tái)的推薦系統(tǒng)大概率遇到過(guò)這種場(chǎng)景用戶(hù)剛點(diǎn)進(jìn)來(lái)時(shí)推薦還算準(zhǔn)但點(diǎn)了三四個(gè)商品之后推薦結(jié)果開(kāi)始跑偏越推越像用戶(hù)很久以前的興趣而不是他此刻正在逛的東西。這個(gè)問(wèn)題在 Session-based Recommendation基于會(huì)話(huà)的推薦里特別典型因?yàn)闀?huì)話(huà)本身很短用戶(hù)沒(méi)有歷史畫(huà)像模型只能靠這一次點(diǎn)擊序列來(lái)猜他下一步想要什么。STAMPShort-Term Attention/Memory Priority Model就是沖著這個(gè)痛點(diǎn)來(lái)的。它的核心主張很直接用戶(hù)的興趣由兩部分組成一部分是這次會(huì)話(huà)里累積出來(lái)的整體興趣general interest另一部分是他最后一次點(diǎn)擊所代表的當(dāng)前興趣current interest。傳統(tǒng)做法用 LSTM 把整個(gè)序列編碼成一個(gè)隱狀態(tài)理論上能記住長(zhǎng)期依賴(lài)但作者在論文里指出LSTM 對(duì)長(zhǎng)會(huì)話(huà)的建模其實(shí)并不夠有效——序列一長(zhǎng)早期信息被稀釋最后那個(gè)隱狀態(tài)未必能準(zhǔn)確反映用戶(hù)現(xiàn)在想要什么。STAMP 的解法是不再只依賴(lài) LSTM 的最終隱狀態(tài)而是顯式地把會(huì)話(huà)平均表示和最后一次點(diǎn)擊表示都拿出來(lái)用一個(gè)注意力網(wǎng)絡(luò)去算每個(gè)歷史 item 對(duì)當(dāng)前興趣的貢獻(xiàn)權(quán)重再加權(quán)求和。這樣既保留了整體興趣的穩(wěn)定性又強(qiáng)化了短期興趣的優(yōu)先級(jí)。適合誰(shuí)看如果你正在做推薦系統(tǒng)、想復(fù)現(xiàn)一個(gè)結(jié)構(gòu)不復(fù)雜但效果扎實(shí)的 baseline或者你已經(jīng)在用 LSTM/GRU 做序列推薦但效果卡住了這篇的配置和驗(yàn)證步驟可以直接拿去跑。我試過(guò)在公開(kāi)數(shù)據(jù)集上把 STAMP 和純 LSTM 版本做對(duì)照差距在短會(huì)話(huà)上尤其明顯。下面從模型結(jié)構(gòu)、配置、訓(xùn)練到排障一步步拆開(kāi)講。2. 環(huán)境與依賴(lài)準(zhǔn)備TaoToken 接入前的模型側(cè)配置在真正寫(xiě) STAMP 之前先把運(yùn)行環(huán)境和依賴(lài)?yán)砬宄?。STAMP 本身是一個(gè)相對(duì)輕量的模型核心就是 embedding 層、注意力層和 MLP 打分層不需要特別重的框架。我一般用 PyTorch 來(lái)復(fù)現(xiàn)因?yàn)樽⒁饬?quán)重的調(diào)試比較直觀(guān)。先建一個(gè)干凈的虛擬環(huán)境把依賴(lài)固定下來(lái)。這里給出一個(gè)可復(fù)制的 requirements 片段路徑按你自己的項(xiàng)目根目錄來(lái)# 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數(shù)據(jù)集方面Session-based Recommendation 最常用的兩個(gè)公開(kāi)數(shù)據(jù)集是 Diginetica 和 Yoochoose現(xiàn)在多叫 RetailRocket 的變體。它們都是會(huì)話(huà)-點(diǎn)擊序列-下一個(gè)點(diǎn)擊的格式。預(yù)處理要做的事很固定按時(shí)間戳切會(huì)話(huà)、過(guò)濾掉長(zhǎng)度小于 2 的會(huì)話(huà)、把 item id 重映射成從 1 開(kāi)始的連續(xù)整數(shù)0 留給 padding。這里有個(gè)容易踩的坑item 重映射一定要在切完訓(xùn)練/測(cè)試集之后統(tǒng)一做否則訓(xùn)練集和測(cè)試集的 id 空間對(duì)不上模型跑起來(lái) loss 會(huì)莫名其妙地不降。我一般寫(xiě)一個(gè)build_vocab函數(shù)先掃全量數(shù)據(jù)建映射表再分別處理各集合。如果你在團(tuán)隊(duì)里做協(xié)作模型代碼和配置建議放到一個(gè)統(tǒng)一的地方管理。我平時(shí)會(huì)把實(shí)驗(yàn)配置、模型權(quán)重路徑、日志目錄都寫(xiě)進(jìn)一個(gè)config.yaml這樣換數(shù)據(jù)集時(shí)只改配置不改代碼。至于模型訓(xùn)練本身STAMP 對(duì)顯存要求不高單卡 8G 就能跑中等規(guī)模數(shù)據(jù)集batch size 設(shè) 128 或 256 都行。環(huán)境準(zhǔn)備好之后下一步就是真正把 STAMP 的結(jié)構(gòu)寫(xiě)出來(lái)。這里要特別注意STAMP 有兩個(gè)版本一個(gè)是 STMP不帶注意力一個(gè)是 STAMP帶注意力。很多人復(fù)現(xiàn)時(shí)直接上 STAMP結(jié)果發(fā)現(xiàn)和論文對(duì)不上其實(shí)是因?yàn)闆](méi)先跑通 STMP 做對(duì)照。建議兩個(gè)都實(shí)現(xiàn)方便驗(yàn)證注意力層到底帶來(lái)了多少提升。3. 可復(fù)制的 STAMP 模型結(jié)構(gòu)與訓(xùn)練配置這一節(jié)是核心直接給可運(yùn)行的模型定義和訓(xùn)練參數(shù)。先看 STAMP 的結(jié)構(gòu)邏輯輸入是一個(gè)會(huì)話(huà)的 item 序列經(jīng)過(guò) embedding 層得到每個(gè) item 的向量然后分兩路一路對(duì)序列做平均得到整體興趣表示一路取最后一個(gè) item 的向量作為當(dāng)前興趣表示接著用注意力機(jī)制計(jì)算每個(gè)歷史 item 對(duì)當(dāng)前興趣的權(quán)重加權(quán)求和得到短期興趣表示最后把整體興趣、當(dāng)前興趣、短期興趣拼接或相加后送入 MLP輸出每個(gè)候選 item 的得分。下面是一個(gè)精簡(jiǎn)但完整的 PyTorch 實(shí)現(xiàn)你可以直接復(fù)制到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, :] # 當(dāng)前興趣 [B, D] avg (emb * mask.unsqueeze(-1)).sum(1) / mask.sum(1, keepdimTrue) # 整體興趣 # 注意力每個(gè)歷史 item 與當(dāng)前興趣的交互 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對(duì)應(yīng)的訓(xùn)練配置我一般寫(xiě)成一個(gè)config.yaml路徑和參數(shù)都固定下來(lái)方便復(fù)現(xiàn)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訓(xùn)練循環(huán)里有個(gè)細(xì)節(jié)要注意STAMP 的損失函數(shù)是標(biāo)準(zhǔn)的交叉熵但負(fù)樣本的構(gòu)造方式會(huì)影響效果。論文里用的是對(duì)每個(gè)正樣本隨機(jī)采樣若干負(fù)樣本的方式我實(shí)測(cè)下來(lái)如果直接用全量 item 做 softmax計(jì)算量大且收斂慢建議先用負(fù)采樣跑通再考慮全量。另外注意力權(quán)重的可視化對(duì)調(diào)試很有幫助。你可以在驗(yàn)證階段把a(bǔ)ttn_weight存下來(lái)看看模型是不是真的把高權(quán)重給了最近幾個(gè)點(diǎn)擊。如果權(quán)重分布很均勻說(shuō)明注意力層沒(méi)學(xué)到東西可能是學(xué)習(xí)率太大或者 embedding 維度太小。4. 驗(yàn)證請(qǐng)求與成功結(jié)果離線(xiàn)評(píng)估怎么跑模型訓(xùn)練完之后必須做離線(xiàn)評(píng)估否則你不知道它到底有沒(méi)有比 baseline 好。Session-based Recommendation 最常用的指標(biāo)是 Recall20 和 MRR20這兩個(gè)指標(biāo)在論文里也是主要對(duì)比項(xiàng)。評(píng)估流程是這樣的對(duì)測(cè)試集里的每個(gè)會(huì)話(huà)取前 n-1 個(gè) item 作為輸入預(yù)測(cè)第 n 個(gè) item模型輸出所有候選 item 的得分取 top-20看真實(shí) 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跑通之后你會(huì)看到類(lèi)似這樣的輸出Epoch 30 | Loss: 2.134 | Recall20: 0.512 | MRR20: 0.221這個(gè)數(shù)字在 Diginetica 上屬于正常范圍。如果你跑出來(lái) Recall20 只有 0.1 左右大概率是數(shù)據(jù)預(yù)處理出了問(wèn)題比如 item id 映射錯(cuò)位或者 padding 沒(méi)處理好。驗(yàn)證階段還有一個(gè)實(shí)用技巧把 STAMP 和 STMP 的評(píng)估結(jié)果放在一起對(duì)比。如果 STAMP 的 Recall20 比 STMP 高 3-5 個(gè)點(diǎn)說(shuō)明注意力層確實(shí)起作用了如果兩者差不多那就要檢查注意力權(quán)重是不是退化了。另外評(píng)估時(shí)要注意測(cè)試集的會(huì)話(huà)長(zhǎng)度分布。如果大部分會(huì)話(huà)都很短比如只有 2-3 個(gè) item那 STAMP 的優(yōu)勢(shì)可能不明顯因?yàn)槎唐谂d趣和整體興趣幾乎重合。這種情況下可以單獨(dú)統(tǒng)計(jì)長(zhǎng)會(huì)話(huà)長(zhǎng)度大于 10上的指標(biāo)更能看出模型差異。5. 本篇常見(jiàn)錯(cuò)排查從 401 到注意力權(quán)重異常復(fù)現(xiàn) STAMP 的過(guò)程中報(bào)錯(cuò)主要集中在幾個(gè)地方。我把自己踩過(guò)的坑列出來(lái)對(duì)照著排查會(huì)快很多。第一個(gè)常見(jiàn)錯(cuò)誤是RuntimeError: expected scalar type Long but found Float。這通常是因?yàn)?embedding 層的輸入要求是整數(shù)類(lèi)型的 item id但你在預(yù)處理時(shí)把 id 轉(zhuǎn)成了 float。解決辦法是檢查seq的數(shù)據(jù)類(lèi)型確保它是torch.long。在 DataLoader 里加一句seq seq.long()就能解決。第二個(gè)是IndexError: index out of range in self。這是 embedding 的經(jīng)典問(wèn)題你的 item id 最大值超過(guò)了num_items。比如你建 vocab 時(shí)統(tǒng)計(jì)的是訓(xùn)練集但測(cè)試集里出現(xiàn)了訓(xùn)練集沒(méi)有的 item。解決辦法是在預(yù)處理階段統(tǒng)一建 vocab或者給未知 item 留一個(gè)專(zhuān)門(mén)的 id。第三個(gè)是注意力權(quán)重全為 0 或者全相等。這通常發(fā)生在 mask 處理不當(dāng)?shù)臅r(shí)候。如果你的mask是 bool 類(lèi)型masked_fill要用mask 0如果是 float 類(lèi)型要確保 padding 位置確實(shí)是 0。我一般會(huì)在 forward 里打印一次attn_weight的均值和方差確認(rèn)它不是一個(gè)常數(shù)。第四個(gè)是 loss 不下降或者震蕩。除了學(xué)習(xí)率太大之外還有一個(gè)容易被忽略的原因負(fù)采樣數(shù)量太少。如果每個(gè)正樣本只采 1 個(gè)負(fù)樣本梯度噪聲會(huì)很大。建議至少采 5-10 個(gè)或者直接用全量 softmax 跑小數(shù)據(jù)集驗(yàn)證。第五個(gè)是評(píng)估指標(biāo)異常低。除了數(shù)據(jù)預(yù)處理問(wèn)題還要檢查評(píng)估時(shí)是不是把 padding 也當(dāng)成了候選 item。正確的做法是在計(jì)算 top-k 時(shí)把 padding 位置的得分設(shè)為負(fù)無(wú)窮。如果你在接入外部服務(wù)做實(shí)驗(yàn)管理時(shí)遇到401 Unauthorized或local proxy failed這類(lèi)報(bào)錯(cuò)通常是鑒權(quán)信息沒(méi)配對(duì)。這時(shí)候可以檢查一下 API Key 是否寫(xiě)進(jìn)了環(huán)境變量以及 Base URL 是否指向了正確的地址。模型側(cè)和平臺(tái)側(cè)的配置要分開(kāi)排查別混在一起調(diào)。6. 語(yǔ)義一致的接入與后續(xù)實(shí)驗(yàn)建議把 STAMP 跑通之后下一步通常是把它接入到實(shí)際的實(shí)驗(yàn)流程里。如果你需要統(tǒng)一管理模型對(duì)話(huà)、API Key 和編碼計(jì)劃可以按下面的路徑操作模型對(duì)話(huà)調(diào)試入口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如果你打算長(zhǎng)期做編碼和 Agent 相關(guān)的實(shí)驗(yàn)Coding Plan 入口在這里https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite控制臺(tái)地址https://taotoken.net/console?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite官網(wǎng)首頁(yè)https://taotoken.net/?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewrite回到模型本身STAMP 之后可以嘗試的改進(jìn)方向有幾個(gè)一是把注意力機(jī)制換成多頭注意力看能不能捕捉更細(xì)的短期興趣二是把 item 的 side information比如類(lèi)別、價(jià)格拼進(jìn) embedding緩解冷啟動(dòng)三是把 STAMP 和 GRU4Rec 做 ensemble取長(zhǎng)補(bǔ)短。這些實(shí)驗(yàn)都可以在現(xiàn)有代碼基礎(chǔ)上改不需要重寫(xiě)整個(gè)框架。最后提醒一句復(fù)現(xiàn)論文模型時(shí)別急著追 SOTA先把 baseline 跑穩(wěn)。STAMP 的價(jià)值不在于它有多復(fù)雜而在于它用很輕的結(jié)構(gòu)把短期興趣優(yōu)先這個(gè)直覺(jué)落到了實(shí)處。你把 STMP 和 STAMP 的對(duì)照實(shí)驗(yàn)做扎實(shí)比盲目堆模塊有用得多。