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

如何统计TensorFlow中自定义Huber损失函数的调用次数?

问题描述

我自定义了Huber损失函数,代码如下:

def my_huber_loss(y_true, y_pred):
    threshold = 1.
    error = y_true - y_pred
    is_small_error = tf.abs(error) <= threshold
    small_error_loss = tf.square(error) / 2
    big_error_loss = threshold * (tf.abs(error) - threshold / 2)
    return tf.where(is_small_error, small_error_loss, big_error_loss)

将其加入model.compile(optimizer='adam', loss=my_huber_loss, metrics=['mae'])后训练正常。

为了统计训练阶段该损失函数的调用次数,我给函数添加了Python计数器:

def my_huber_loss(y_true, y_pred):
    threshold = 1.
    error = y_true - y_pred
    is_small_error = tf.abs(error) <= threshold
    small_error_loss = tf.square(error) / 2
    big_error_loss = threshold * (tf.abs(error) - threshold / 2)
    my_huber_loss.counter +=1 # 新增计数器更新
    return tf.where(is_small_error, small_error_loss, big_error_loss)

my_huber_loss.counter = 0 # 初始化计数器

但训练结束后执行print(my_huber_loss.counter)输出结果为3,明显不符合实际调用次数。另外我在损失函数中添加tf.print("--- Called Loss ---"),能看到训练过程中函数被多次调用,说明计数器统计失效。

问题原因

TensorFlow在训练时默认使用计算图执行模式:当你把Python写的损失函数传入model.compile时,TensorFlow会将这个函数转化为计算图(即进行函数追踪)。此时你添加的my_huber_loss.counter +=1是Python层面的操作,只会在图构建阶段执行几次(比如初始化、验证输入形状等),而实际训练时,模型是在计算图中执行,不会再触发Python层面的计数器累加。

而tf.print是TensorFlow的图内操作,每次计算图运行(每个batch计算损失时)都会执行,所以能看到大量打印,但Python计数器只在图构建时更新了几次,最终得到3的结果。

正确统计调用次数的方法

方法1:使用TensorFlow变量作为计数器

使用tf.Variable替代Python变量,因为tf.Variable属于计算图的一部分,每次图执行都会更新它的值:

import tensorflow as tf

def my_huber_loss(y_true, y_pred):
    threshold = 1.
    error = y_true - y_pred
    is_small_error = tf.abs(error) <= threshold
    small_error_loss = tf.square(error) / 2
    big_error_loss = threshold * (tf.abs(error) - threshold / 2)
    # 更新TensorFlow变量计数器
    my_huber_loss.counter.assign_add(1)
    return tf.where(is_small_error, small_error_loss, big_error_loss)

# 初始化TensorFlow变量计数器
my_huber_loss.counter = tf.Variable(0, dtype=tf.int32)

训练结束后,通过.numpy()读取计数器的值:

print(my_huber_loss.counter.numpy())

方法2:自定义损失类(更规范)

继承tf.keras.losses.Loss类,在类中维护计数器,这种方式更符合Keras的设计规范,也更易维护:

import tensorflow as tf

class HuberLoss(tf.keras.losses.Loss):
    def __init__(self, threshold=1.0, name='huber_loss'):
        super().__init__(name=name)
        self.threshold = threshold
        # 初始化计数器
        self.counter = tf.Variable(0, dtype=tf.int32)
    
    def call(self, y_true, y_pred):
        error = y_true - y_pred
        is_small_error = tf.abs(error) <= self.threshold
        small_error_loss = tf.square(error) / 2
        big_error_loss = self.threshold * (tf.abs(error) - self.threshold / 2)
        # 更新计数器
        self.counter.assign_add(1)
        return tf.where(is_small_error, small_error_loss, big_error_loss)

# 使用损失类实例
huber_loss = HuberLoss(threshold=1.0)
model.compile(optimizer='adam', loss=huber_loss, metrics=['mae'])

# 训练后查看计数器
print(huber_loss.counter.numpy())

内容的提问来源于stack exchange,提问作者Theo Deep

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 15:30:44