艺术领域协同过滤模型训练遇InvalidArgumentError求解决方案
解决协同过滤模型训练时的InvalidArgumentError问题
问题核心
训练面向艺术领域的协同过滤模型时,触发如下错误:
InvalidArgumentError: Graph execution error:
Node: 'CollaborativeFiltering/xusers_emb/embedding_lookup'
indices[28,0] = 1000 is not in [0, 1000)
该错误本质是用户ID的索引值超出了嵌入层(Embedding Layer)的输入范围:嵌入层定义的有效输入区间为[0, 1000)(即0到999),但数据中出现了值为1000的用户ID,导致嵌入查找操作失败。
解决方案
1. 修正嵌入层的input_dim参数
嵌入层的input_dim需要设置为最大用户ID + 1(因为索引从0开始计数)。如果你的数据集最大用户ID是1000,那么input_dim应设为1001,确保覆盖0到1000的所有索引值:
# 用户嵌入层示例代码 user_input = Input(shape=(1,), name='user_input') user_embedding = Embedding(input_dim=1001, output_dim=64, input_length=1)(user_input)
同时检查艺术品嵌入层的input_dim,确保与艺术品ID的最大取值匹配。
2. 将用户/艺术品ID映射为连续索引
如果原始数据中的用户ID并非从0开始的连续整数(存在间断或起始值不为0),需重新映射为连续索引,避免出现超出范围的ID:
import pandas as pd # 基于训练+验证集的所有用户ID创建统一映射 all_user_ids = pd.concat([train["user_id"], val["user_id"]]).unique() user_id_mapping = {original_id: idx for idx, original_id in enumerate(all_user_ids)} # 应用映射到训练集和验证集 train["user_id"] = train["user_id"].map(user_id_mapping) val["user_id"] = val["user_id"].map(user_id_mapping) # 对艺术品ID执行相同的映射操作 all_art_ids = pd.concat([train["art_id"], val["art_id"]]).unique() art_id_mapping = {original_id: idx for idx, original_id in enumerate(all_art_ids)} train["art_id"] = train["art_id"].map(art_id_mapping) val["art_id"] = val["art_id"].map(art_id_mapping)
注意:必须基于全量数据(训练+验证)创建映射,避免验证集中出现训练集未覆盖的ID,导致索引越界。
3. 调整数据集拆分顺序
如果使用validation_split=0.3自动拆分验证集,需先完成ID映射,再启动模型训练。若先拆分再映射,会导致训练集与验证集的ID映射规则不一致,触发索引错误。
4. 修正训练代码的缩进问题
你的训练代码存在缩进错误,需调整为同一层级:
training = model.fit(x=[train["user_id"], train["art_id"]], y=train["y"], epochs=100, batch_size=128, shuffle=True, verbose=0, validation_split=0.3) model = training.model utils_plot_keras_training(training)
内容的提问来源于stack exchange,提问作者Andrii
相关产品推荐
相关产品推荐

