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.pymain()函数中,首先通过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-1m

3. 执行数据处理

运行shell脚本启动数据处理流程:

bash run_ml-1m.sh

4. 检查输出结果

处理完成后,在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),仅供参考