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

Keras中模型摘要正常却出现形状不兼容错误

Keras有状态LSTM训练时形状不兼容问题排查

问题背景

搭建了带跳连接的有状态LSTM模型,模型摘要显示各层输出形状匹配,但训练时触发形状不兼容错误,将stateful设为False后模型可正常运行,自定义损失函数可能影响反向传播的形状兼容性。

模型代码

# Since we are predicting a value for every timestep, we set return_sequences=True
input = Input(batch_shape=ip_shape)
mLSTM = LSTM(units=32, return_sequences=True, stateful=True)(input)
mDense = Dense(units=32, activation='linear')(input)
mSkip = Add()([mLSTM, mDense])

mSkip = Dense(units=1, activation='linear')(mSkip)
model = Model(input, mSkip)

adam = Adam(learning_rate=0.01)
model.compile(optimizer=adam, loss=total_loss)
model.summary()

模型摘要

Model: "model_3"
_______________________________________________________________________________________________
 Layer (type)                Output Shape                 Param #   Connected to                  
===============================================================================================
 input_5 (InputLayer)        [(104, 22050, 1)]            0         []                            
                                                                                                    
 lstm_4 (LSTM)               (104, 22050, 32)             4352      ['input_5[0][0]']             
                                                                                                    
 dense_5 (Dense)             (104, 22050, 32)             64        ['input_5[0][0]']             
                                                                                                    
 add_3 (Add)                 (104, 22050, 32)             0         ['lstm_4[0][0]',               
                                                                     'dense_5[0][0]']             
                                                                                                    
 dense_6 (Dense)             (104, 22050, 1)              33        ['add_3[0][0]']                
                                                                                                    
===============================================================================================
Total params: 4449 (17.38 KB)
Trainable params: 4449 (17.38 KB)
Non-trainable params: 0 (0.00 Byte)
_______________________________________________________________________________________________

自定义损失函数

def total_loss(y_true, y_pred):
    ratio = 0.5
    dc_loss = math_ops.pow(math_ops.subtract(math_ops.mean(y_true, 0), math_ops.mean(y_pred, 0)), 2)
    dc_loss = math_ops.mean(dc_loss, axis=-1)
    dc_energy = math_ops.mean(math_ops.pow(y_true, 2), axis=-1) + 0.00001
    dc_loss = math_ops.div(dc_loss, dc_energy)

    esr_loss = math_ops.squared_difference(y_pred, y_true) 
    esr_loss = math_ops.mean(esr_loss, axis=-1)
    esr_energy = math_ops.mean(math_ops.pow(y_true, 2), axis=-1) + 0.00001
    esr_loss = math_ops.div(esr_loss, esr_energy)

    return (ratio)*dc_loss + (1-ratio)*esr_loss

训练错误信息

InvalidArgumentError: Graph execution error:

...

Incompatible shapes: [104,22050,32] vs. [32,22050,1]
     [[{{node gradient_tape/total_loss/BroadcastGradientArgs}}]] [Op:__inference_train_function_9604]

问题原因

  1. 有状态LSTM的batch维度处理特性
    当stateful=True时,Keras会固定保留每一批输入的LSTM隐藏状态,要求输入的batch_shape固定,且梯度计算流程中,batch维度(轴0)的张量处理逻辑与无状态模式不同——梯度计算时会对batch维度的状态进行特殊追踪,导致张量的形状维度顺序或结构发生隐式调整。

  2. 自定义损失函数的轴操作冲突
    损失函数中math_ops.mean(y_true, 0)直接对**batch维度(轴0)**求均值,会将原本形状为[104,22050,1]的y_true压缩为[22050,1],后续与其他张量进行广播运算时,在有状态LSTM的梯度计算语境下,原本兼容的形状出现错位,触发形状不匹配错误。

  3. 无状态模式正常的原因
    当stateful=False时,模型不会保留批次状态,batch维度的处理遵循常规张量运算逻辑,损失函数的轴操作不会干扰梯度计算的张量形状兼容性,因此不会触发错误。

解决方法

  • 调整损失函数的轴计算逻辑:避免直接对batch维度(轴0)求均值,改为针对时间步或特征维度计算,或者明确指定多轴均值以保证形状兼容性。例如,若需计算全局均值,可改为math_ops.mean(y_true, axis=[0,1,2])(根据实际需求调整轴参数),确保损失函数输出的形状与梯度计算预期一致。
  • 验证损失函数输出形状:自定义损失函数应返回标量或与batch维度匹配的张量(每个样本对应一个损失值),可在损失函数中添加形状检查语句,确认前向传播时的形状符合预期。
  • 规范有状态LSTM的训练流程:确保训练时按固定批次顺序输入数据,并在每个epoch开始时调用model.reset_states()重置状态,避免状态残留间接影响张量形状处理。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 09:02:49