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

TensorFlow训练最后迭代报错:Reshape输入维度不匹配问题排查

解决TensorFlow DMRL模型训练时的Reshape维度不匹配错误

问题描述

运行基于TensorFlow的DMRL模型时,训练至最后迭代出现错误:

Input to reshape is a tensor with 28 values, but the requested shape has 128

错误关联节点为:gradient_tape/model/ItemEmbedding/embedding_lookup/Reshape_1。怀疑问题出在FactorInteractionLayer类中,但无法定位具体位置。

相关模型代码、FactorInteractionLayer代码及训练数据维度如下:

模型核心代码

def DMRL(n_users, n_items, embed_dim, n_factors):
    assert embed_dim % n_factors == 0, "embed_dim must be divisible by n_factors"
    
    user_input = Input(shape=(1,), dtype='int32', name='UserInput')
    user_embedding = Embedding(n_users, embed_dim, name='UserEmbedding')(user_input)
    users = Flatten(name='UserFlatten')(user_embedding)
    
    item_input = Input(shape=(1,), dtype='int32', name='ItemInput')
    item_embedding = Embedding(n_items, embed_dim, name='ItemEmbedding')(item_input)
    items = Flatten(name='ItemFlatten')(item_embedding)

    textual_input = Input(shape=(768,), name='TextualInput')
    textual_mlp = Modal_MLP(embed_dim, 4)(textual_input)

    visual_input = Input(shape=(4096,), name='VisualInput')
    visual_mlp = Modal_MLP(embed_dim, 4)(visual_input)
    
    user_factor_embedding = tf.split(users, n_factors, 1)
    item_factor_embedding = tf.split(items, n_factors, 1)
    textual_factor_embedding = tf.split(textual_mlp, n_factors, 1)
    visual_factor_embedding = tf.split(visual_mlp, n_factors, 1)

    #Disentangled Representation Learning
    cor_loss = CorrelationLossLayer(n_factors)([visual_factor_embedding, textual_factor_embedding, user_factor_embedding, item_factor_embedding])

    # Factor Interaction
    rating = FactorInteractionLayer(n_factors)([user_factor_embedding, item_factor_embedding, textual_factor_embedding, visual_factor_embedding])
    rating = Flatten()(rating)

    output = Dense(1, activation='linear')(rating)
    
    model = Model(inputs=[user_input, item_input, textual_input, visual_input], outputs=[output, cor_loss])

    model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate = 0.01),
                  loss=[mse_loss, cor_loss_dummy],  
                  loss_weights=[1.0, 0.2])
    return model

def mse_loss(y_true, y_pred):
    return tf.keras.losses.MeanSquaredError()(y_true, y_pred)

def cor_loss_dummy(y_true, y_pred):
    return y_pred

emebd_dim = 128
n_factors = 2
batch_size = emebd_dim//n_factors

model.fit([train_user, train_item, train_text, train_image],
                [train_y, tf.zeros_like(train_y)],  
                batch_size=batch_size,
                epochs=50,
                callbacks=[es],
                validation_split=0.1)

训练数据维度

  • train_user shape: (275232,)
  • train_item shape: (275232,)
  • train_text shape: (275232, 768)
  • train_image shape: (275232, 4096)
  • train_y.shape: (275232,)

FactorInteractionLayer原代码

class FactorInteractionLayer(Layer):
    def __init__(self, n_factors):
        super(FactorInteractionLayer, self).__init__()
        self.n_factors = n_factors
        self.h = Dense(3, activation='tanh')
        self.user_sig = Dense(1, activation='sigmoid')
        self.text_sig = Dense(1, activation='sigmoid')
        self.visual_sig = Dense(1, activation='sigmoid')
        self.attention_layer = Dense(3, activation='softmax')

    def call(self, inputs):
        user_embedding, item_embedding, text_embedding, visual_embedding = inputs[0], inputs[1], inputs[2], inputs[3]

        output = 0
        for i in range(self.n_factors):
            user_emb, item_emb, text_emb, visual_emb = user_embedding[i], item_embedding[i], text_embedding[i], visual_embedding[i]

            user_item_text_visual = Concatenate()([user_emb, item_emb, text_emb, visual_emb])
            h = self.h(user_item_text_visual)
            attention_weights = self.attention_layer(h)

            user_item = tf.matmul(user_emb, item_emb)
            user_item = self.user_sig(user_item)
            user_text = tf.matmul(user_emb, text_emb)
            user_text = self.text_sig(user_text)
            user_visual = tf.matmul(user_emb, visual_emb)
            user_visual = self.visual_sig(user_visual)

            user_item_importance = attention_weights[:, 0]
            user_text_importance = attention_weights[:, 1]
            user_visual_importance = attention_weights[:, 2]

            user_item_imp = tf.tensordot(user_item, user_item_importance, axes=0)
            user_text_imp = tf.tensordot(user_text, user_text_importance, axes=0)
            user_visual_imp = tf.tensordot(user_visual, user_visual_importance, axes=0)

            sum_concat = Concatenate()([user_item_imp, user_text_imp, user_visual_imp])

            output += tf.reduce_sum(sum_concat)
        output = tf.expand_dims(output, -1)
        return output

错误根源定位

错误的核心是**FactorInteractionLayer中维度计算逻辑混乱**,导致模型输出与期望维度不匹配,反向传播时触发Embedding层的Reshape错误:

  1. 向量交互计算错误:使用tf.matmul计算用户-物品等交互时,输入是(batch_size, 64)的向量,直接matmul会得到(batch_size, batch_size)的矩阵,完全偏离了单样本交互的预期,导致后续维度爆炸。
  2. 注意力权重结合错误:tf.tensordot(user_item, user_item_importance, axes=0)会将(batch_size,1)的交互值与(batch_size,)的权重生成(batch_size, batch_size)的张量,进一步扩大维度混乱。
  3. 输出累加逻辑错误:每次迭代对sum_concat做tf.reduce_sum,最终得到的是全局标量而非每个样本的输出,导致模型输出维度为(1,),与训练数据的(batch_size,1)标签维度不匹配,反向传播时触发Embedding层的梯度维度错误。

解决方案

针对上述问题,修正FactorInteractionLayer的核心逻辑:

  • 用元素相乘+求和替代tf.matmul,计算单样本的向量交互值(保持(batch_size,1)维度)。
  • 用元素乘法替代tensordot,将注意力权重与交互值按样本维度结合。
  • 保留batch维度进行累加,最终输出每个样本的交互结果,匹配训练数据维度。

修正后的FactorInteractionLayer代码

class FactorInteractionLayer(Layer):
    def __init__(self, n_factors):
        super(FactorInteractionLayer, self).__init__()
        self.n_factors = n_factors
        self.h = Dense(3, activation='tanh')
        self.user_sig = Dense(1, activation='sigmoid')
        self.text_sig = Dense(1, activation='sigmoid')
        self.visual_sig = Dense(1, activation='sigmoid')
        self.attention_layer = Dense(3, activation='softmax')

    def call(self, inputs):
        user_embedding, item_embedding, text_embedding, visual_embedding = inputs[0], inputs[1], inputs[2], inputs[3]
        
        # 初始化output为batch维度的0张量
        batch_size = tf.shape(user_embedding[0])[0]
        output = tf.zeros((batch_size, 3), dtype=tf.float32)
        
        for i in range(self.n_factors):
            user_emb, item_emb, text_emb, visual_emb = user_embedding[i], item_embedding[i], text_embedding[i], visual_embedding[i]

            user_item_text_visual = Concatenate()([user_emb, item_emb, text_emb, visual_emb])
            h = self.h(user_item_text_visual)
            attention_weights = self.attention_layer(h)

            # 修正:用元素相乘+求和计算向量点积,得到单样本交互值
            user_item = tf.reduce_sum(tf.multiply(user_emb, item_emb), axis=1, keepdims=True)
            user_item = self.user_sig(user_item)
            user_text = tf.reduce_sum(tf.multiply(user_emb, text_emb), axis=1, keepdims=True)
            user_text = self.text_sig(user_text)
            user_visual = tf.reduce_sum(tf.multiply(user_emb, visual_emb), axis=1, keepdims=True)
            user_visual = self.visual_sig(user_visual)

            # 修正:扩展权重维度后做元素乘法,保持batch维度
            user_item_importance = tf.expand_dims(attention_weights[:, 0], -1)
            user_text_importance = tf.expand_dims(attention_weights[:, 1], -1)
            user_visual_importance = tf.expand_dims(attention_weights[:, 2], -1)

            user_item_imp = user_item * user_item_importance
            user_text_imp = user_text * user_text_importance
            user_visual_imp = user_visual * user_visual_importance

            sum_concat = Concatenate(axis=1)([user_item_imp, user_text_imp, user_visual_imp])
            # 按batch维度累加
            output += sum_concat
        
        return output

同时注意原模型中rating = Flatten()(rating)可以保留,因为修正后的输出是(batch_size,3),Flatten后是(batch_size,3),再经过Dense(1)得到最终的(batch_size,1)预测值,完全匹配标签维度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 07:28:10