TensorFlow TextVectorization二维输入秩报错解决方案
问题根因
经过listwise采样生成训练数据后,输入的movie_title张量形状为(batch_size, 5),即每个batch中每个用户对应5个候选电影组成的排序列表,属于二维字符串张量。但TextVectorization层仅支持两种合法输入:一维字符串张量(形状为(n,)),或最后一维长度为1的字符串张量,直接传入二维张量就会触发你遇到的报错。
直接使用tf.flatten()展平无效的原因是:展平操作会丢失「每个用户对应5个候选」的列表维度信息,后续做listwise损失计算时无法对齐用户、候选、评分的对应关系,会触发新的维度不匹配错误。
修复步骤
核心处理逻辑:传入embedding层前临时展平张量满足层输入要求,特征计算完成后恢复原有列表维度,不需要改动上游数据采样逻辑。
- 第一步:修改
MovieModel的call方法,适配二维列表输入
def call(self, titles): # 保存输入原始形状 [batch_size, 候选列表长度] original_shape = tf.shape(titles) # 展平为一维张量,满足StringLookup和TextVectorization的输入要求 flattened_titles = tf.reshape(titles, [-1]) # 分别计算电影ID嵌入、电影标题文本嵌入 id_embedding = self.title_embedding(flattened_titles) text_embedding = self.title_text_embedding(flattened_titles) # 拼接特征后恢复列表维度,输出形状 [batch_size, 候选列表长度, 嵌入维度] combined_embedding = tf.concat([id_embedding, text_embedding], axis=1) return tf.reshape(combined_embedding, [original_shape[0], original_shape[1], -1])
- 第二步:修复其余阻断运行的代码错误
- 变量名错误:
MovielensModel.call中定义了query_embeddings存储用户嵌入,后续代码错误引用了未定义的user_embeddings,统一替换为query_embeddings即可。 - 用户嵌入维度适配:用户嵌入原始形状为
[batch_size, 嵌入维度],和候选电影嵌入拼接前,需要扩展维度并按候选列表长度复制,对齐形状为[batch_size, 候选列表长度, 嵌入维度],示例代码:
# 扩展用户嵌入维度,复制为和候选列表等长 list_len = tf.shape(movie_embeddings)[1] query_embeddings_tiled = tf.tile( tf.expand_dims(query_embeddings, 1), [1, list_len, 1] ) # 再拼接用户和电影嵌入计算评分 rating_predictions = self.rating_model( tf.concat([query_embeddings_tiled, movie_embeddings], axis=2) ) # 把预测结果形状压缩为[batch_size, 候选列表长度],匹配标签形状 rating_predictions = tf.squeeze(rating_predictions, axis=-1)- 任务定义修正:你在
compute_loss中调用了未初始化的self.retrieval_task,如果要实现listwise排序损失,直接在__init__中给Ranking任务传入列表损失即可,不需要额外定义检索任务:
self.rating_task = tfrs.tasks.Ranking( loss=tfr.keras.losses.ListMLELoss(), metrics=[tfr.keras.metrics.NDCGMetric(name="ndcg@5", topn=5)], )- 损失计算逻辑修正:
compute_loss中不需要计算retrieval_loss,直接返回rating_loss即可,同时要保证传入的预测值和标签形状完全一致。
- 变量名错误:
内容的提问来源于stack exchange,提问作者sakeesh
相关产品推荐
相关产品推荐

