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

TensorFlow Recommenders训练报错Shape需rank2实为rank3,需基于make_csv_dataset修复

错误根源

  • ID编码逻辑错误:你的user_id是UUID格式字符串、channel_id包含英文字母,直接调用tf.strings.to_number无法正常转换为合法数值,且没有做字符串ID到整数索引的映射,Embedding层无法直接处理这类非连续输入。
  • 张量维度不匹配:tf.data.experimental.make_csv_dataset本身已经按照你设置的batch_size=1返回带批次维度的张量,后续你又对候选集channels调用了.batch(1),等于叠加了两层批次维度,导致候选embedding从期望的(样本数, 64)变成了(1,1,64)的3维张量,触发了报错中的维度等级不匹配问题。
  • 候选集配置错误:make_csv_dataset默认开启shuffle,会导致候选集每次加载顺序随机,FactorizedTopK指标计算时无法获取稳定的候选集,也会触发计算错误。

适配要求的完整可运行代码(保留tf.data.experimental.make_csv_dataset流式加载逻辑)

from typing import Dict, Text
import pandas as pd
from pathlib import Path

import tensorflow as tf 
import tensorflow_recommenders as tfrs

# 生成测试数据集
df_interactions = pd.DataFrame({
    'user_id': [
        '00001446-da5f-4d17', 
        '00001446-da5f-4d17',
        '00005ab5-c9e0-4b05-',
        '00005ab5-c9e0-4b05-',
        '000093dd-1a11-4600', 
        '000093dd-1a11-4600',
        '00009b34-65b5-42c1', 
        '0000ae32-4a91-4bcd',
        '0000ae32-4a91-4bcd',
        '0000ae32-4a91-4bcd'
    ], 
    'channel_id': [
        '1', '2', 'A56',
        '3', 'B72', '2', 
        'M63', '2', '5', 'A56'
    ]
})
df_interactions.to_csv('experiment_interactions.csv', index=False)

df_channels = pd.DataFrame({
    'channel_id': [
        '1', '2', '3', '5', 'A56', 'B72', 'M63' 
    ],
    'channel_name': [
        'Popular', 
        'Best',
        'Highest Rated',
        'Large Following',
        'Nice', 
        'Retro',
        'Modern'
    ]
})
df_channels.to_csv('experiment_channels.csv', index=False)

# --------------------------
# 流式加载CSV核心逻辑修改
# --------------------------
# 加载交互数据,临时batch设为1后续unbatch,方便单样本处理
interactions = tf.data.experimental.make_csv_dataset(
    file_pattern='experiment_interactions.csv', 
    column_defaults=[tf.string, tf.string], 
    batch_size=1,
    header=True,
    shuffle_seed=42
).unbatch() # 去掉多余的批次维度,返回单样本的数据集

# 加载候选频道数据,关闭shuffle保证候选集顺序稳定
channels = tf.data.experimental.make_csv_dataset(
    file_pattern='experiment_channels.csv', 
    column_defaults=[tf.string, tf.string], 
    batch_size=1,
    header=True,
    shuffle=False
).unbatch().map(lambda x: x["channel_id"]) # 仅保留channel_id字段,去掉多余批次维度

# --------------------------
# 字符串ID转整数索引逻辑
# --------------------------
# 构建用户ID词汇表
user_ids = interactions.map(lambda x: x["user_id"])
user_vocab = tf.keras.layers.StringLookup(mask_token=None)
user_vocab.adapt(user_ids)

# 构建频道ID词汇表
channel_vocab = tf.keras.layers.StringLookup(mask_token=None)
channel_vocab.adapt(channels)

# 预处理交互数据,转ID为整数索引
interactions = interactions.map(lambda x: {
    "user_id": user_vocab(x["user_id"]),
    "channel_id": channel_vocab(x["channel_id"])
})

# --------------------------
# 模型定义
# --------------------------
class Model(tfrs.Model):
  def __init__(self):
    super().__init__()
    # 用户表示层,输入维度为词汇表大小+1(StringLookup默认从1开始索引,0为OOV)
    self.user_model = tf.keras.layers.Embedding(
        input_dim=user_vocab.vocabulary_size(), output_dim=64)
    # 频道表示层
    self.item_model = tf.keras.layers.Embedding(
        input_dim=channel_vocab.vocabulary_size(), output_dim=64)
    # 检索任务,候选集直接映射为embedding,不要额外加batch
    self.task = tfrs.tasks.Retrieval(
        metrics=tfrs.metrics.FactorizedTopK(
            candidates=channels.map(channel_vocab).map(self.item_model).batch(128)
        )
    )

  def compute_loss(self, features: Dict[Text, tf.Tensor], training=False) -> tf.Tensor:
    user_embeddings = self.user_model(features["user_id"])
    channel_embeddings = self.item_model(features["channel_id"])
    return self.task(user_embeddings, channel_embeddings)

# 训练配置
model = Model()
model.compile(optimizer=tf.keras.optimizers.Adagrad(0.5))
# 训练时统一做批次处理
model.fit(interactions.batch(8), epochs=5)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 17:24:07