tflite_model_maker训练推荐模型时报错MapDataset无gen_dataset属性
问题解决方法
错误原因
报错'MapDataset' object has no attribute 'gen_dataset'是因为tflite_model_maker的推荐任务不接受原生TensorFlow的tf.data.Dataset类型作为输入,要求传入封装好的recommendation.DataLoader类实例,该内置类自带训练所需的gen_dataset方法。
解决步骤
步骤1:预处理原始数据为序列样本
推荐任务的训练样本需要包含用户历史交互上下文序列和待预测的下一个交互物品标签,首先从原始CSV生成符合要求的样本格式:
import pandas as pd from tflite_model_maker import recommendation # 读取原始数据集 df = pd.read_csv('./dataset.csv') # 按用户ID分组,生成每个用户的历史交互序列(有交互时间的话建议先按时间排序) user_interaction_seq = df.groupby('user_id')['post_ids'].apply(list).reset_index() # 滑动窗口生成训练样本 train_sample_list = [] for seq in user_interaction_seq['post_ids']: # 交互数不足2的用户无法生成有效样本,直接过滤 if len(seq) < 2: continue # 取前i个交互作为上下文,第i+1个作为标签 for i in range(1, len(seq)): train_sample_list.append({ "context": seq[:i], "label": seq[i] })
步骤2:生成物品ID词表
需要提前生成所有物品ID的词表,预留特殊标记位:
# 提取所有唯一的post_id all_post_ids = df['post_ids'].unique().tolist() vocab_save_path = "./post_id_vocab.txt" # 写入词表,前两位固定为填充标记<PAD>和未知标记<OOV> with open(vocab_save_path, "w", encoding="utf-8") as f: f.write("<PAD>\n") f.write("<OOV>\n") for pid in all_post_ids: f.write(f"{pid}\n")
步骤3:初始化模型配置和数据集加载器
# 初始化模型配置,传入词表路径、最大历史长度参数 model_spec = recommendation.ModelSpec( model_name="recommendation_bow", vocab_file=vocab_save_path, max_history_length=10 ) # 加载自定义数据集为DataLoader实例 train_data = recommendation.DataLoader.from_list(train_sample_list, model_spec)
步骤4:启动训练
model_x = recommendation.create( train_data, model_spec=model_spec, batch_size=16, steps_per_epoch=10000, epochs=1, learning_rate=0.1, gradient_clip_norm=1.0, shuffle=True, do_train=True )
注意事项
- 若原始数据已经提前拆分好了上下文和标签,可直接跳过滑动窗口生成步骤,按要求格式拼装
train_sample_list即可 - 词表顺序不可随意调整,
<PAD>必须对应索引0,<OOV>必须对应索引1 - 如需加入用户ID特征,可额外配置用户词表,同时在训练样本中添加
user_id字段
内容的提问来源于stack exchange,提问作者Chrome Phantom
相关产品推荐
相关产品推荐

