TensorFlow model.fit底层逻辑:训练与验证损失计算方式问询
KeyboardInterrupt Traceback (most recent call last)
**编辑补充**:附上我的自定义损失函数代码: ```python def global_loss(predicted_probability, predicted_energy, time, true_probability, true_energy): #from the input data, get everything to have a shape of [None, span] None is a placeholder in graph execution for batch_size predicted_probability = tf.cast(tf.squeeze(predicted_probability, axis = -1), tf.float32) predicted_energy = tf.cast(tf.squeeze(predicted_energy, axis = -1), tf.float32) t = tf.cast(time, tf.float32) true_energy = tf.cast(true_energy, tf.float32) true_probability = tf.cast(true_probability, tf.float32) #predicted photon count loss predicted_photon_count = tf.math.reduce_sum(predicted_probability) true_photon_count = tf.math.reduce_sum(true_probability) predicted_photon_count_loss = tf.math.square(predicted_photon_count - true_photon_count) #predicted photon energy loss predicted_photon_energy = tf.math.reduce_sum(tf.math.multiply(predicted_probability, predicted_energy)) true_photon_energy = tf.math.reduce_sum(true_energy) predicted_photon_energy_loss = tf.math.square(predicted_photon_energy - true_photon_energy) #predicted energy variance loss predicted_energy_variance = tf.math.reduce_sum(tf.math.multiply(tf.math.square(tf.math.subtract(predicted_energy, predicted_photon_energy)), predicted_probability)) true_energy_variance = tf.math.reduce_sum(tf.math.square(tf.math.subtract(true_energy, true_photon_energy))) predicted_energy_variance_loss = tf.square(predicted_energy_variance - true_energy_variance) #predicted photon incidence loss predicted_photon_incidence = tf.math.reduce_sum(tf.math.multiply(t, predicted_probability)) true_photon_incidence = tf.math.reduce_sum(tf.math.multiply(t,true_probability)) predicted_photon_incidence_loss = tf.math.square(predicted_photon_incidence - true_photon_incidence) #predicted time variance loss predicted_time_variance = tf.math.reduce_sum(tf.math.multiply(tf.math.square(tf.math.subtract(t, predicted_photon_incidence)), predicted_probability)) true_time_variance = tf.math.reduce_sum(tf.math.multiply(tf.math.square(tf.math.subtract(t, true_photon_incidence)), true_probability)) predicted_time_variance_loss = tf.math.square(predicted_time_variance - true_time_variance) #weighting photon_count_weight = tf.constant(1e2, dtype=tf.float32) * predicted_photon_count_loss photon_energy_weight = tf.constant(5e1, dtype=tf.float32) * predicted_photon_energy_loss energy_variance_weight = tf.constant(1, dtype=tf.float32) * predicted_energy_variance_loss photon_incidence_weight = tf.constant(1, dtype=tf.float32) * predicted_photon_incidence_loss time_variance_weight = tf.constant(1e2, dtype=tf.float32) * predicted_time_variance_loss #returning loss differently based on whether a photon is in the batch global_loss = tf.cond( pred = tf.math.equal(true_photon_count, tf.constant(0., dtype = tf.float32)), true_fn = lambda: photon_count_weight, false_fn = lambda: photon_count_weight + photon_energy_weight + energy_variance_weight + photon_incidence_weight + time_variance_weight ) return global_loss
解答
当你自定义train_step和test_step时,model.fit()对每个batch返回的损失值默认采用加权平均的方式计算epoch级别的损失,权重为每个batch的样本数量。
结合你的代码细节来看:
- 你的
global_loss函数中,所有损失项都是通过tf.reduce_sum计算的整个batch的总损失,而非单样本损失的平均值。比如predicted_photon_count = tf.math.reduce_sum(predicted_probability),最终返回的是当前batch所有样本的损失总和。 model.fit()会收集每个batch返回的总损失值,然后根据每个batch包含的样本数做加权平均,得到最终显示的epoch损失(即"I II II L"和val_loss)。如果所有batch的样本数相同,这个结果就等同于所有batch损失的简单平均值;如果存在样本数不同的batch(比如最后一个batch),则会按样本数比例加权。
另外注意你代码中的一处错误:在train_step更新energy模型权重时,目标权重写错了,应该用self.energy.trainable_weights而非重复使用self.probability.trainable_weights,修正后代码如下:
grads = tape.gradient(loss, self.energy.trainable_weights) self.optimizer.apply_gradients( zip(grads, self.energy.trainable_weights) )
内容的提问来源于stack exchange,提问作者Jezen_Volk
相关产品推荐
相关产品推荐

