You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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])
  • 第二步:修复其余阻断运行的代码错误
    1. 变量名错误:MovielensModel.call中定义了query_embeddings存储用户嵌入,后续代码错误引用了未定义的user_embeddings,统一替换为query_embeddings即可。
    2. 用户嵌入维度适配:用户嵌入原始形状为[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)
    
    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)],
    )
    
    1. 损失计算逻辑修正:compute_loss中不需要计算retrieval_loss,直接返回rating_loss即可,同时要保证传入的预测值和标签形状完全一致。

内容的提问来源于stack exchange,提问作者sakeesh

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.27 02:03:25