Keras调用train_on_batch时出现梯度为None的错误求助
排查自定义损失函数梯度为None的问题
咱们先拆解下你遇到的问题:虽然keras.losses.binary_crossentropy是可微分的,但你的模型架构和损失定义方式让Keras无法正确追踪梯度传递的路径,才导致了这个错误。
错误根源分析
- 损失层的定义方式不对:你把损失计算封装成了Lambda层作为模型的输出,然后在
compile时给这个输出指定了lambda y_true, y_pred: y_pred的损失函数。这相当于告诉Keras“这个输出的损失就是它自己”,但Keras无法从这个已经计算好的损失值反向传播到模型的LSTM、Dense等参数层——因为这个Lambda层的输出和模型参数之间的梯度链路被切断了。 - train_on_batch的目标参数不匹配:你的模型有两个输出(
y_pred和loss_out),但调用train_on_batch时只传了noise作为目标,这不仅不符合输出数量,而且loss_out作为损失计算结果,根本不需要对应的目标值。
修正后的解决方案
正确的做法是用add_loss方法在模型内部添加自定义损失,让Keras自动追踪梯度链路,不需要把损失作为模型输出。下面是修正后的代码:
import numpy as np from keras import backend as K import keras from keras.models import Model from keras.layers import TimeDistributed, Dense, Dropout, LSTM, Input def my_loss(input_y, input_y_pred): # 直接用模型输入的input_y和input_y_pred计算二分类交叉熵损失 return keras.losses.binary_crossentropy(input_y, input_y_pred) def generator2(): input_noise = Input(name='input_noise', shape=(40, 38), dtype='float32') input_y = Input(name='input_y', shape=(1,), dtype='float32') input_y_pred = Input(name='input_y_pred', shape=(1,), dtype='float32') # 模型主体结构保持不变 lstm1 = LSTM(256, return_sequences=True)(input_noise) drop = Dropout(0.2)(lstm1) lstm2 = LSTM(256, return_sequences=True)(drop) y_pred = TimeDistributed(Dense(38, activation='softmax'))(lstm2) # 关键:用add_loss将自定义损失添加到模型中,Keras会自动处理梯度 custom_loss = my_loss(input_y, input_y_pred) model = Model(inputs=[input_noise, input_y, input_y_pred], outputs=[y_pred]) model.add_loss(custom_loss) # 编译时无需指定loss参数,因为已通过add_loss添加 model.compile(optimizer='adam') return model g2 = generator2() noise = np.random.uniform(0,1,size=[10,40,38]) # train_on_batch只需传入输入,目标参数传None(若只有自定义损失) g2.train_on_batch([noise, np.ones(10), np.zeros(10)], None)
额外说明
如果你的y_pred本身还有对应的任务损失(比如序列分类损失),可以同时添加多个损失:
# 假设新增一个真实标签输入用于y_pred的分类损失 y_true_input = Input(name='y_true', shape=(40, 38), dtype='float32') # 给y_pred添加分类损失 classification_loss = keras.losses.categorical_crossentropy(y_true_input, y_pred) # 同时添加两个损失 model.add_loss(custom_loss + classification_loss) # 此时train_on_batch需要传入所有输入 g2.train_on_batch([noise, np.ones(10), np.zeros(10), y_true_data], None)
这样调整后,Keras就能正确追踪所有参数的梯度,不会再出现“梯度为None”的错误。
内容的提问来源于stack exchange,提问作者Thiago Alves
相关产品推荐
相关产品推荐

