如何统计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

