BERT4Rec数据处理实战:从原始数据到TFRecord的高效转换
BERT4Rec数据处理实战:从原始数据到TFRecord的高效转换
【免费下载链接】BERT4RecBERT4Rec: Sequential Recommendation with Bidirectional Encoder Representations from Transformer项目地址: https://gitcode.com/gh_mirrors/be/BERT4Rec
BERT4Rec作为基于Transformer的序列推荐模型,其数据处理流程是模型性能的关键环节。本文将详细介绍如何使用BERT4Rec项目中的工具,将原始用户-物品交互数据高效转换为TFRecord格式,为模型训练提供高质量输入。
数据处理核心工具概述
BERT4Rec的数据处理主要依赖于两个核心脚本,它们共同构成了从原始数据到训练数据的完整流水线:
- gen_data.py:负责数据预处理、序列构建和TFRecord文件生成
- vocab.py:处理词汇表构建,将物品ID转换为模型可识别的索引
这两个脚本配合工作,实现了从原始文本数据到模型输入的全自动化转换,支持多种数据集和配置参数。
原始数据格式解析
BERT4Rec支持的原始数据存储在项目的data目录下,如:
- data/ml-1m.txt:MovieLens-1M数据集
- data/beauty.txt:亚马逊Beauty数据集
- data/steam.txt:Steam游戏数据集
这些文件采用简单的文本格式,每行代表一个用户的物品交互序列,格式为用户ID 物品ID1 物品ID2 ... 物品IDn,物品ID按交互时间排序。例如:
1 101 205 310 ... 2 502 108 42 ...数据处理完整流程
1. 数据加载与划分
在gen_data.py的main()函数中,首先通过data_partition()函数加载原始数据并划分为训练集、验证集和测试集:
dataset = data_partition(output_dir+dataset_name+'.txt') [user_train, user_valid, user_test, usernum, itemnum] = dataset默认情况下,验证集会合并到训练集中,形成最终的训练数据:
# put validate into train for u in user_train: if u in user_valid: user_train[u].extend(user_valid[u])2. 词汇表构建
词汇表构建是将物品ID映射为整数索引的关键步骤,由FreqVocab类实现(位于vocab.py):
vocab = FreqVocab(user_test_data)词汇表会自动为特殊标记(如[CLS]、[MASK]、[PAD])预留索引,并根据物品出现频率分配索引值,确保高频物品有较小的索引值。
3. 训练实例生成
create_training_instances()函数是数据处理的核心,它将用户交互序列转换为模型可训练的实例:
instances = create_training_instances( data, max_seq_length, dupe_factor, short_seq_prob, masked_lm_prob, max_predictions_per_seq, rng, vocab, mask_prob, prop_sliding_window, force_last=False)该过程包含以下关键步骤:
- 序列截断与滑动窗口:当序列长度超过
max_seq_length时,使用滑动窗口切分长序列 - 数据增强:通过
dupe_factor参数控制数据重复次数,每次重复应用不同的掩码策略 - 掩码语言模型(MLM)预处理:随机掩盖序列中的物品,用于模型训练
4. TFRecord文件生成
最后,write_instance_to_example_files()函数将训练实例写入TFRecord文件:
writers.append(tf.python_io.TFRecordWriter(output_file))TFRecord格式的优势在于:
- 高效的磁盘I/O性能
- 支持分布式训练
- 内置压缩机制节省存储空间
生成的TFRecord文件默认保存在data目录下,命名格式为{dataset_name}{version_id}.train.tfrecord。
关键参数配置
通过命令行参数可以灵活控制数据处理过程,主要参数包括:
| 参数 | 作用 | 默认值 |
|---|---|---|
max_seq_length | 序列最大长度 | 200 |
masked_lm_prob | 掩码概率 | 0.15 |
dupe_factor | 数据重复次数 | 10 |
prop_sliding_window | 滑动窗口步长比例 | 0.1 |
dataset_name | 数据集名称 | ml-1m |
实际使用时,可以通过修改run_ml-1m.sh等脚本中的参数来适应不同的数据集和训练需求。
实战操作步骤
1. 准备原始数据
将原始数据文件(如ml-1m.txt)放置在data目录下,确保格式符合要求。
2. 配置参数
修改对应的shell脚本,如处理MovieLens-1M数据集时编辑run_ml-1m.sh:
--max_seq_length=128 \ --masked_lm_prob=0.15 \ --dupe_factor=10 \ --dataset_name=ml-1m3. 执行数据处理
运行shell脚本启动数据处理流程:
bash run_ml-1m.sh4. 检查输出结果
处理完成后,在data目录下会生成:
- TFRecord文件:如ml-1mdefault.train.tfrecord
- 词汇表文件:如ml-1mdefault.vocab
- 历史数据文件:如ml-1mdefault.his
常见问题解决
数据格式错误
如果原始数据格式不符合要求,会导致data_partition()函数解析失败。解决方法:
- 确保每行格式为"用户ID 物品ID1 物品ID2 ..."
- 检查是否存在空行或格式不一致的行
内存占用过高
处理大型数据集(如ml-20m.txt)时可能出现内存问题:
- 减小
max_seq_length参数 - 降低
dupe_factor值 - 分批次处理数据
TFRecord文件过大
可以通过修改代码将输出文件分割为多个小文件,提高并行处理效率:
# 在write_instance_to_example_files函数中 output_files = [output_file + ".part" + str(i) for i in range(num_shards)]总结
BERT4Rec的数据处理流程通过gen_data.py和vocab.py实现了从原始交互数据到TFRecord格式的完整转换。该流程具有高度的灵活性和可配置性,能够适应不同规模和类型的推荐系统数据集。通过合理调整参数,可以为模型训练提供最优的输入数据,从而提升序列推荐性能。
掌握这一数据处理流程,不仅能帮助你更好地使用BERT4Rec模型,也能为其他序列推荐模型的数据预处理提供参考思路。
【免费下载链接】BERT4RecBERT4Rec: Sequential Recommendation with Bidirectional Encoder Representations from Transformer项目地址: https://gitcode.com/gh_mirrors/be/BERT4Rec
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考