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错误:
- 向量交互计算错误:使用
tf.matmul计算用户-物品等交互时,输入是(batch_size, 64)的向量,直接matmul会得到(batch_size, batch_size)的矩阵,完全偏离了单样本交互的预期,导致后续维度爆炸。 - 注意力权重结合错误:
tf.tensordot(user_item, user_item_importance, axes=0)会将(batch_size,1)的交互值与(batch_size,)的权重生成(batch_size, batch_size)的张量,进一步扩大维度混乱。 - 输出累加逻辑错误:每次迭代对
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
相关产品推荐
相关产品推荐

