ARTICLE DETAIL

资讯详情

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

DREAM模型实战:基于循环神经网络的下一篮子推荐系统解析

DREAM模型实战:基于循环神经网络的下一篮子推荐系统解析 简介面向电商个性化推荐场景的完整Python深度学习源码包聚焦“下一篮子推荐”任务适用于推荐算法入门者及数据科学从业者深入研究用户时序行为建模与商品补全预测。资源共14个文件包体约20KB核心包含6个Python脚本完整覆盖数据清洗、RNN模型构建、模型训练与测试评估全流程另含3个JSON样例数据文件训练、测试及验证便于直接运行验证。目前已有61人学习下载。项目内附带依赖清单、配置文件、CI集成测试及说明文档目录结构清晰完整可帮助快速复现DREAM等经典序列推荐模型。读者可借此掌握从购物篮数据预处理、嵌入表示、时序特征提取到离线指标评估的完整闭环并深入理解协同过滤思想与深度学习结合解决用户“下一单可能买什么”的实际工程思路。1. 基于神经网络的下一篮子推荐先搞清它和普通推荐的区别在电商场景里最常见的推荐是“看了又看”和“买过还买”但真正影响客单价的往往是这样一个问题用户上周买了牛奶和面包这周购物车里会出现什么下一篮子推荐Next-Basket Recommendation就是专门回答这个问题的任务它和普通点击率预估最大的区别在于输入不是一个孤立行为而是一串按时间排序的购物篮。这个仓库里的 DREAM 模型用循环神经网络把历史篮子序列编码成用户的动态表示再预测下一个篮子的商品组成。对做推荐系统、研究序列建模的工程师来说这是一份能直接跑通、能改参、能换数据的完整代码不是那种只给半截网络定义的示例。2. 项目结构与数据管线从三个 JSON 样本到可训练的序列2.1 仓库文件布局先看懂每个文件负责哪一段打开压缩包第一眼不是模型代码而是 data 目录下那三个 JSON 样本和一堆零散的 .py 文件。很多人拿到项目就直奔 rnn_model.py结果连数据格式都没摸清跑起来报错报得莫名其妙。我拿推荐项目的第一步永远是先确认三件事输入数据长什么样、模型吃什么格式、评估在哪个指标上做。这个仓库的布局其实很清楚我把文件职责拆成一张表文件职责关键注意点data/train_sample.json训练集样本模型只在它上面学data/validation_sample.json验证集样本用来挑超参、提前停data/test_sample.json测试集样本最终评估指标只看这里utils/dataprocess.py把原始 JSON 构造成用户篮子序列输出格式决定模型输入utils/data_helpers.py词典构建、padding、批处理训练循环直接调用DREAM/config.py所有超参数集中管理调参先改这里DREAM/rnn_model.pyDREAM 网络本体GRU 注意力DREAM/train.py训练入口命令行启动DREAM/test.py测试入口输出候选和指标这个布局最值得夸的一点是数据处理和模型定义完全分开。很多开源推荐项目把数据 pipeline 写死在训练脚本里想换数据集就得通读几百行代码。这里 dataprocess 和 data_helpers 独立出来意味着只要保证 dataprocess 的输出格式不变换数据源的成本就控制在几个函数以内。2.2 数据预处理购物篮序列是怎么变成模型输入的下一篮子推荐的数据最原始形式通常是订单表一行是用户某次下单的商品集合。仓库里的 JSON 样本按用户拆分每个用户有一串按时间排序的篮子每个篮子是一个商品 ID 集合。这种变长嵌套结构没法直接喂给 RNN必须转成定长张量dataprocess.py 干的就是这件事。第一步是把原始记录聚合为“每个用户的历史篮子列表”# utils/dataprocess.py 核心流程结构保留细节按实际实现 def build_user_sequences(raw_data): user_sequences {} for record in raw_data: uid record[user_id] basket [item[product_id] for item in record[basket]] user_sequences.setdefault(uid, []).append(basket) for uid in user_sequences: user_sequences[uid].sort(keylambda b: record[date]) return user_sequences这里有两个关键点。一是 basket 内只保留 product_id数量、价格、折扣这些信息在 DREAM 模型里不参与因为下一篮子推荐的核心是“买过什么集合”不是“买了多少”二是排序必须依据交易时间戳而不是文件里的出现顺序否则序列语义就全乱了。你要是换了业务数据这两点最容易翻车。第二步是切分输入输出对。取用户前 n-1 个篮子作为输入序列第 n 个篮子作为预测目标def generate_train_pairs(user_sequences, max_basket_num10): X, y [], [] for uid, baskets in user_sequences.items(): for i in range(1, len(baskets)): input_seq baskets[max(0, i - max_basket_num): i] target_basket baskets[i] X.append(input_seq) y.append(target_basket) return X, y这里的max_basket_num是窗口长度。用户历史可能有一百个篮子但 RNN 不可能从头到尾记住所有信息取最近 10 个篮子就够了。窗口太大会引入噪声窗口太小丢失长期依赖这个参数要和数据处理一起调。目标值 y 也不是一个商品 ID而是一个“篮子”——所以不能直接做单标签分类要映射成多热向量用多个二分类输出来表达“这个商品在不在下一个篮子里”。第三步就轮到 data_helpers.py 了。它做两件事先统计所有商品 ID 构建词典再把每个篮子映射成商品索引向量最后在 batch 内 padding 成等长。padding 的坑我后面专门讲这里先记住一个原则padding 位置在计算 attention 时一定要 mask 掉否则模型会把“空位”当成真实商品学进去。2.3 训练/测试数据划分验证集不是让你“再学一点”用的仓库给的是三个 JSON 文件不少新手会直接合并 validation 和 train 一起训练这是最容易后悔的行为之一。验证集的唯一用途是过程中判断收敛和过拟合。DREAM 这类 RNN 模型后期训练 loss 下降很慢但验证集的 RecallN 可能已经在往下掉——这就是过拟合信号该停了。你要是把验证集混进训练集等于提前把答案告诉模型最后测试指标全是虚高上线就被打脸。还要注意一个细节这三个 sample 文件是从完整数据集里随机抽出来的不是按时间切分的。也就是说同一个用户完全可能同时在 train_sample 和 test_sample 里出现。学术实验能这么干因为 DREAM 不依赖用户 ID embedding只依赖行为序列模式但你要是把这套代码直接搬上业务数据必须改成按时间切分——用前 80% 时间的交易训练后 20% 时间交易测试否则会时间穿越用未来的数据训练过去的测试线上指标没有参考价值。注意sample 文件是随机抽样而不是时间切分直接拿指标上线前做流量评估数字会虚高一到两个点。生产环境一定要自己按时间重新切分。3. DREAM 模型RNN 怎么把一堆购物篮编码成用户表示3.1 config.py 参数全解哪些参数动一个全盘受影响DREAM 这个名字拆开看就是 Dynamic REcurrent Attention Model三个词对应三个设计决策用户表示随时间动态变化、用循环网络建模变化、用注意力选择关键历史。你要改模型行为先动 config.py 里的这些旋钮参数常见取值影响hidden_size100 / 128RNN 隐层维度代表用户表示的表达能力embedding_dim50 / 64商品嵌入维度太小商品向量分不开num_layers1 / 2RNN 层数加深不等于变好dropout0.2 ~ 0.5防过拟合序列任务里尤其关键learning_rate0.001 ~ 0.01Adam 下的常见区间batch_size32 / 64结合显存和样本量max_basket_num10 / 15输入窗口长度top_k5 / 10评估时取前 K 个商品算指标参数之间不是独立的。embedding_dim 太小不同商品映射到相近向量推荐结果会偏向热门商品冷门商品永远出不来hidden_size 太大在小数据上很容易过拟合训练集 loss 一路下降但验证集指标原地踏步。我一般先用最小配置跑通流程再按两倍粒度逐级放大一次只动一个变量不然你根本不知道指标变化是哪个参数引起的。优化器方面仓库默认用的是 Adam少数版本会看到 SGD 的配置。我的经验是序列推荐任务别换 SGDAdam 对梯度的自适应缩放能让训练稳定很多尤其在你频繁调学习率的时候。学习率调度可以试试 ReduceLROnPlateau验证集 Recall 连续三个 epoch 不涨就把学习率减半这个操作往往比调结构更快看到收益。3.2 网络结构拆解篮子内平均序列上用 GRU历史间加注意力DREAM 的 forward 分成三层嵌入层、GRU 编码层、注意力融合层。下面这段是核心结构的简化实现# DREAM/rnn_model.py简化自仓库实现 import torch import torch.nn as nn class DREAM(nn.Module): def __init__(self, vocab_size, embedding_dim, hidden_size, num_layers, dropout, max_basket_num): super().__init__() self.embedding nn.Embedding(vocab_size, embedding_dim) self.rnn nn.GRU(embedding_dim, hidden_size, num_layers, batch_firstTrue, dropoutdropout) self.attn nn.Linear(hidden_size, 1) def forward(self, basket_seq, maskNone): # basket_seq: [batch, seq_len, basket_size] 商品ID basket_emb self.embedding(basket_seq) # [batch, seq, basket, emb] basket_emb basket_emb.mean(dim2) # [batch, seq, emb] rnn_out, _ self.rnn(basket_emb) # [batch, seq, hidden] score self.attn(rnn_out).squeeze(-1) # [batch, seq] if mask is not None: score score.masked_fill(mask 0, -1e9) attn_weight torch.softmax(score, dim1) user_rep (rnn_out * attn_weight.unsqueeze(-1)).sum(dim1) return user_rep这个结构里有三个设计细节值得细品。第一每个篮子内多个商品先做平均 pooling再接 GRU。购物篮本质上是无序集合商品在篮子内没有先后关系强行按顺序编码反而制造噪声。第二RNN 选 GRU 而不是 LSTM两者在序列建模能力上接近但 GRU 参数少、收敛快在样本量不充裕的推荐场景里更稳。第三attention 分数在 softmax 之前必须把 padding 位置置为负无穷这样权重才会均匀分配到真实篮子上——不 mask 的话padding 位置的零向量会和真实篮子竞争权重模型直接学偏。GRU 的隐藏状态在这里扮演的是“用户当前偏好”的角色。每输入一个新篮子隐藏状态就被更新一次更新幅度由 GRU 的门控机制决定——如果新篮子内容和历史偏好一致更新就小状态稳定如果出现了一个全新类目的商品门控会让状态大幅跳跃。这种动态性正是 DREAM 名字里 Dynamic 的由来用户表示不是静态的 embedding而是跟着行为序列一路演化的状态向量。3.3 输出层与损失函数为什么不直接用交叉熵GRU 出来的是用户表示向量离“预测下一个篮子”还差最后一步把 user_rep 映射回商品空间得到每个商品被选中的概率。实现上就是一个线性层输出维度等于词典大小。但损失函数不能直接套多分类交叉熵因为预测目标是一个商品集合不是一个商品 ID。正确的做法是把输出当成多个独立的二分类任务对每个商品用 BCEWithLogitsLosslogits self.fc(user_rep) # [batch, vocab_size] loss nn.BCEWithLogitsLoss()(logits, target_basket_onehot)target 是“这个商品是否出现在下一篮子中”的多热向量。用 BCE 而不是交叉熵意味着模型要学会的是一组独立的伯努利分布这样可以同时处理购物篮大小变化的问题有的用户下一篮子只有一件商品有的有五件BCE 不要求归一化模型能自由调整输出的商品数量。这也解释了为什么 DREAM 评估时用 RecallN 而不是精确率——候选商品多、真实篮小精确率容易被稀释召回更能反映“该中的有没有中”。除了 RecallN如果你要写技术报告或和同行对比建议把 NDCGN 也一起算上。Recall 只关心命中没命中NDCG 还关心命中的商品在 top_N 里的排名位置。两个模型可能在 Recall10 上打平但 NDCG 能区分出谁把正确答案排得更靠前后者对线上用户体验更敏感。计算 NDCG 时注意对每个用户单独归一化然后用平均算数平均汇总别直接用全局排序算那样会把长尾用户淹没。4. 训练与测试跑通 train.py 和 test.py 的完整操作4.1 环境准备版本匹配是第一步requirements.txt 一般会列出核心依赖。我建议直接用 conda 建干净环境别在系统全局 Python 里乱装这个项目依赖深度学习框架和数据处理库版本混了很容易出现互相踩踏的情况conda create -n nbr python3.8 conda activate nbr pip install -r requirements.txt如果你机器上已经装了新版 PyTorch2.x要注意旧项目里可能用了被移除的接口比如torch.autograd.Variable。装完依赖先跑一句检查命令然后grep -rn Variable DREAM/看代码里有没有要改的地方python -c import torch, json, numpy; print(torch.__version__, torch.cuda.is_available()) grep -rn Variable DREAM/有Variable的地方改成torch.tensor就行。这一步十分钟能省掉后面一小时排错。GPU 不是必须的这个数据规模用 CPU 也能跑但训练会慢不少有条件还是让 PyTorch 走 CUDA。4.2 训练阶段参数怎么传日志怎么读数据放好后进入 DREAM 目录启动训练cd DREAM python train.py \ --data_dir ../data \ --epochs 20 \ --batch_size 32 \ --lr 0.001 \ --hidden_size 100 \ --embedding_dim 50 \ --save_dir ./checkpoints这里有几个参数需要解释。--data_dir指向的是放 train_sample.json、validation_sample.json 的目录train.py 会自动按文件名找--epochs别一上来给 50先用 20 看 loss 掉到多少平台期再决定--save_dir跑之前确认目录存在有些版本的代码没有os.makedirs(exist_okTrue)目录不存在直接 IOError 崩溃。训练日志每个 epoch 会打印 loss 和验证集指标。看日志时别只盯 lossloss 下降不代表模型变好。我经历过训练 loss 从 5.2 降到 1.1验证集 Recall10 却卡在 0.12 不动的情况这种基本是 embedding 学崩了或者注意力权重全堆在最后一个时间步。正确做法是每两三个 epoch 同时记录 loss 和 Recall看到 loss 降但 Recall 不再涨就手动提前停。提示保存 checkpoint 时把 config 里的超参一起序列化到同一个文件下次加载模型不会因为参数不一致出现维度对不上的问题。4.3 测试阶段从模型输出到评估指标训练完执行测试python test.py \ --data_dir ../data \ --checkpoint ./checkpoints/model_best.pt \ --top_k 10test.py 会加载保存的模型对 test_sample.json 里的每个用户计算商品概率分布输出 top_k 个商品 ID。输出文件每行一般是user_id, item_id1 item_id2 ...拿这个可以直接和真实篮子算指标。如果你要自己写评估脚本核心计算长这样def compute_recall_at_k(preds, true_basket, k10): # preds: 按概率降序排列的商品ID列表 if not true_basket: return 0.0 hit len(set(preds[:k]) set(true_basket)) return hit / len(set(true_basket))这里有个非常容易被忽略的问题如果 test.py 输出结果全是热门商品先别怀疑代码去检查评估时有没有过滤用户已经买过的商品。很多序列推荐模型有隐藏 bug——预测时没有把历史篮子里出现过的商品屏蔽掉结果模型学到的是“用户买过洗发水下次还推荐洗发水”训练指标好看线上毫无意义。严谨的做法是在 test.py 里加一步过滤把用户历史篮子中的商品 ID 从候选集中剔除再取 top_k。5. 避坑指南跑 DREAM 项目最常见的五个翻车点5.1 数据加载与预处理阶段的两个坑坑一JSON 字段对不上导致 KeyError 崩溃现象json.load()执行正常一进入 dataprocess 就报KeyError: product_id。原因换了自己的数据后字段名不一致有的叫item_id有的叫product_id有的 basket 字段嵌套层级不一样。仓库自带的 sample 文件字段是统一的但业务数据几乎不会刚好同名。解决加载后先打印第一条记录看结构再做一层字段归一化不要直接在核心代码里硬编码字段名。def normalize_record(record): return { user_id: record.get(user_id) or record.get(uid), basket: record.get(basket) or record.get(items, []), }坑二训练和测试的词典不一致导致索引越界现象训练一切正常测试时报IndexError: index out of range in self。原因vocab 是在训练时构建的测试数据里出现了训练集中没见过的商品 IDembedding 层索引不到。解决把训练阶段构建好的item2idx和idx2item保存成 JSON 文件测试时加载同一个词典而不是在 test.py 里重新构建。见过没见过的商品统一映射成UNK。5.2 训练与评估阶段的三个坑坑三loss 在训练中段突然变成 NaN现象前几个 epoch 正常第五个 epoch 左右 loss 突然变成 NaN训练卡死。原因最常见的是学习率过大导致梯度爆炸。其次是前面提到的 padding 没 masksoftmax 分母里出现 0 或极小数反向传播梯度直接变成无穷。解决把 learning_rate 降到 0.0005 以下同时在 attention 计算时显式处理 padding mask。两件事一起做能杜绝绝大多数 NaN 问题。score self.attn(rnn_out).squeeze(-1) score score.masked_fill(mask 0, -1e9) attn_weight torch.softmax(score, dim1)坑四测试集 Recall5 恒为 0现象模型训练指标正常但一跑 test.pyRecall 全部为零。原因测试集里混入了只有一个篮子的用户。一个篮子没法构成“历史→预测”对模型在预测时完全没有历史依据输出等于瞎猜。另一个可能是真实篮子和候选集完全没有商品交集说明词典映射出了问题。解决在数据处理阶段筛掉历史篮子数少于 2 的用户并跑一遍商品 ID 交集检查确认测试集的商品都被训练集词典覆盖。坑五训练时间异常长一个 epoch 要十几分钟现象batch_size 设了 32数据量才几万条但每个 epoch 耗时夸张。原因没有做长度分桶。长短序列混在一个 batch 里RNN 按 batch 内最长序列展开大量时间花在 padding 空位上。解决训练前按用户历史篮子数排序分桶让同一个 batch 内的序列长度接近速度能提升一倍以上。这个优化对线上训练尤其重要模型迭代速度直接决定你能试多少组参数。6. 调优与验证从 RecallN 到业务效果的最后一公里基础版本跑通之后真正的工程活才开始。先别急着加模型复杂度有三件成本最低的改进收益比扩网络大得多。第一把篮子内平均 pooling 换成注意力 pooling。购物篮有主次之分用户买牛奶可能是常规行为但给新生儿买奶粉的那次购物基本主导了接下来的购物内容。用注意力对不同商品加权比平均把所有商品当成同等重要的是更贴近真实行为的。实现就是在 basket 内商品维度上再叠加一层 attention改动不超过二十行。第二纠正数据划分方式。随机切分会高估模型好几个点我一般用最近四周的交易做测试之前全部做训练。改完你会看到 Recall 下降一两个点这是正常的说明之前的指标确实有水。线上评估要用时间切分后的指标说话不要拿随机切分的数字给老板看。第三做一个最朴素的 baseline 对比。用“用户最近一个篮子里出现频率最高的商品”作为推荐结果算 RecallN。如果 DREAM 比这个简单 baseline 的 Recall 高不到 10%问题不在模型在数据本身——商品序列太稀疏或者篮子间关联太弱再好的模型也学不出来。这个 baseline 十行代码就能写完def last_basket_baseline(user_sequences): preds {} for uid, baskets in user_sequences.items(): counter collections.Counter(baskets[-1]) preds[uid] [item for item, _ in counter.most_common(10)] return preds最后从指标回到业务RecallN 高了不代表线上转化就好。如果推荐结果全是用户不可能再买的大件消费品虽然从“同一用户”角度看命中但从购买周期看是错的。我会在评估脚本额外输出一个“复购占比”——推荐商品里有多少在用户历史篮子中出现过。占比过高说明模型在学重复购买对新品类无探索能力占比过低说明模型在乱探索忘了用户基本需求。这个平衡点没有标准答案要看业务想要拉复购还是做新品渗透。我第一次跑这套代码就是在“测试集 Recall 全零”这个坑里耗了一整晚最后发现是未过滤单篮子用户。从那以后我每次跑序列推荐项目都会先做数据统计看每用户篮子数的分布直方图、商品覆盖率和序列长度中位数这三张图看完再进模型。这个习惯帮我避开了至少五六个白费功夫的夜晚。这份 Next-Basket-Recommendation 的原始代码包值得你下载下来照着跑一遍里面的数据处理和模型定义比任何论文附录都实在。希望这份拆解也能帮你少走几段弯路。本文还有配套的精品资源点击获取
返回列表