RecBole训练BERT4Rec/SASRec效果极差,求问题排查与解决
问题:RecBole实现BERT4Rec/SASRec在ml-1m上效果远低于论文预期
我基于RecBole实现BERT4Rec与SASRec模型,对比论文结果时发现:用ml-1m数据集训练效果极差,受算力限制训练轮数未超50轮,但20轮左右指标就趋于平稳,ndcg@10仅0.0523,远低于论文宣称的0.48+。调整过大量超参数(包括对齐论文参数),但指标最高没突破0.07,同时训练损失极高。
配置文件
#Enviroment settings gpu_id: 0 log_wandb: true train_neg_sample_args: ~ learning_rate: 0.001 weight_decay: 0.005 mask_ratio: 0.2 hidden_size: 64 data_path: C:\user\MyScripts\dataset load_col: inter: [user_id, item_id, rating, timestamp] item: [item_id, movie_title, release_year, class] user: [user_id, age, gender, occupation, zip_code] threshold: {'rating': 3}
运行代码
from recbole.quick_start import run_recbole run_recbole(model='BERT4Rec', dataset='ml-1m', config_file_list = ['configExample.yaml'])
训练日志
Train 40: 100%|█████████████████████████| 48/48 [00:07<00:00, 6.19it/s, GPU RAM: 2.11 G/10.00 G] 17 Apr 13:38 INFO epoch 40 training [time: 7.75s, train loss: 280.5661] Evaluate : 100%|███████████████████████████| 1/1 [00:00<00:00, 15.09it/s, GPU RAM: 2.11 G/10.00 G] 17 Apr 13:38 INFO epoch 40 evaluating [time: 0.07s, valid_score: 0.034600] 17 Apr 13:38 INFO valid result: recall@10 : 0.1113 mrr@10 : 0.0346 ndcg@10 : 0.0523 hit@10 : 0.1113 precision@10 : 0.0111 17 Apr 13:38 INFO Finished training, best eval result in epoch 29
解决方案
1. 数据处理修正
- 调整正样本阈值:论文中ml-1m采用隐式反馈,将评分>=4视为正样本,当前
threshold: {'rating': 3}会引入大量弱正样本,干扰模型学习,建议改为threshold: {'rating': 4}。 - 配置负采样策略:当前
train_neg_sample_args: ~未开启负采样,BERT4Rec依赖负样本训练,建议设置:
对齐论文的单负样本采样策略。train_neg_sample_args: {'sample_num': 1, 'distribution': 'uniform', 'dynamic': False, 'candidate_num': 0}
2. 超参数与模型配置对齐
- 优化学习率与权重衰减:论文中BERT4Rec使用
1e-4的学习率,当前0.001过大导致训练不稳定、损失过高,建议改为learning_rate: 0.0001;权重衰减0.005偏大,改为weight_decay: 0.0或1e-4。 - 补充核心模型参数:论文中BERT4Rec为2层2注意力头,序列长度设为50,需在配置中添加:
num_hidden_layers: 2 num_attention_heads: 2 max_seq_length: 50 - 配置优化器参数:如需修改beta1/beta2,可在配置中明确指定:
optimizer: Adam adam_beta1: 0.9 adam_beta2: 0.999
3. 训练策略调整
- 增加有效训练轮数:论文训练轮数超100轮,当前50轮可能未充分收敛。若算力有限,可增大
train_batch_size(如设为256),加快单轮训练速度,从而完成更多轮次训练。 - 确认损失函数:确保使用RecBole默认的
BERT4RecLoss(交叉熵损失),避免自定义损失函数导致偏差。
内容的提问来源于stack exchange,提问作者TBD Foxix
相关产品推荐
相关产品推荐

