【发布时间】:2020-12-03 00:58:25
【问题描述】:
我正在学习 tensorflow(2.3) 中的 keras API。在 tensorflow 网站上的这个 guide 中,我找到了一个自定义损失函数的示例:
def custom_mean_squared_error(y_true, y_pred):
return tf.math.reduce_mean(tf.square(y_true - y_pred))
此自定义损失函数中的reduce_mean 函数将返回一个标量。
这样定义损失函数是否正确?据我所知,y_true 和y_pred 形状的第一个维度是批量大小。我认为损失函数应该为批次中的每个样本返回损失值。所以损失函数应该给出一个形状为(batch_size,)的数组。但是上面的函数为整个批次给出了一个单一的值。
也许上面的例子是错误的?谁能帮我解决这个问题?
附言为什么我认为损失函数应该返回一个数组而不是单个值?
我阅读了Model类的源代码。当您向Model.compile() 方法提供损失函数(请注意它是函数,而不是损失类)时,该损失函数用于构造LossesContainer 对象,存储在Model.compiled_loss。这个传递给LossesContainer类的构造函数的损失函数再次用于构造一个LossFunctionWrapper对象,该对象存储在LossesContainer._losses中。
根据LossFunctionWrapper类的源码,通过LossFunctionWrapper.__call__()方法(继承自Loss类)计算一个训练batch的整体损失值,即返回单个损失值整个批次。 但是LossFunctionWrapper.__call__() 首先调用LossFunctionWrapper.call() 方法来获取训练批次中每个样本的损失数组。然后这些损失最终被平均以获得整个批次的单个损失值。在LossFunctionWrapper.call() 方法中调用了提供给Model.compile() 方法的损失函数。
这就是为什么我认为自定义损失函数应该返回一系列损失,而不是单个标量值。此外,如果我们为Model.compile() 方法编写一个自定义Loss 类,我们自定义Loss 类的call() 方法也应该返回一个数组,而不是一个信号值。
我在 github 上开了一个issue。已确认需要自定义损失函数才能为每个样本返回一个损失值。该示例需要更新以反映这一点。
【问题讨论】:
标签: tensorflow machine-learning keras tensorflow2.0 loss-function