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

TensorFlow模型拟合报错:logits与labels第一维度不匹配

LSTM+BERT模型训练维度不匹配问题排查

模型代码

# Define the model
def build_lstm_model():
    input_ids = tf.keras.layers.Input(shape=(128,), dtype=tf.int32, name='input_ids')
    
    # BERT embedding layer
    bert_model = TFBertModel.from_pretrained('bert-base-uncased')
    bert_output = bert_model(input_ids)[1]  # using the pooled output

    # LSTM layer
    lstm_output = tf.keras.layers.LSTM(64)(tf.expand_dims(bert_output, axis=1))  # Expand the dimensions for LSTM

    # Output layer
    output = tf.keras.layers.Dense(1, activation='softmax')(lstm_output)

    model = tf.keras.Model(inputs=input_ids, outputs=output)
    model.compile(optimizer='adam',
                  loss='sparse_categorical_crossentropy',
                  metrics=['accuracy'])
    return model

# Train the model
model = build_lstm_model()
#model.summary()

model.fit(train_dataset, validation_data=val_dataset, epochs=3)

报错信息

logits and labels must have the same first dimension, got logits shape [8,1] and labels shape [1024]

补充说明

  • train_dataset和val_dataset均为shape=(None,128)的TensorFlow数据集
  • 模型结构输出(model.summary()):
Model: "model"
_________________________________________________________________
 Layer (type)                Output Shape              Param #   
=================================================================
 input_ids (InputLayer)      [(None, 128)]             0         
                                                                 
 tf_bert_model (TFBertModel  TFBaseModelOutputWithPo   109482240 
 )                           olingAndCrossAttentions             
                             (last_hidden_state=(Non             
                             e, 128, 768),                       
                              pooler_output=(None, 7             
                             68),                                
                              past_key_values=None,              
                             hidden_states=None, att             
                             entions=None, cross_att             
                             entions=None)                       
                                                                 
 tf.expand_dims (TFOpLambda  (None, 1, 768)            0         
 )                                                               
                                                                 
 lstm (LSTM)                 (None, 64)                213248    
                                                                 
 dense (Dense)               (None, 1)                 65        
                                                                 
=================================================================
Total params: 109695553 (418.46 MB)
Trainable params: 109695553 (418.46 MB)
Non-trainable params: 0 (0.00 Byte)
_________________________________________________________________

已尝试的解决方案

  • 扁平化LSTM层
lstm_output = tf.keras.layers.Flatten()(lstm_output)
  • 修改LSTM层和Dense层的输出形状
  • 将损失函数改为categorical_crossentropy
  • 将输出层激活函数改为sigmoid

请求帮助解决该维度匹配问题。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 19:12:48