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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 01:15:35