tf.function装饰层call方法时add_loss触发超出作用域错误
TensorFlow中add_loss方法使用问题解答
核心结论
TensorFlow计算图中完全支持通过add_loss方法手动添加自定义损失项,这是Keras自定义层、自定义模型提供的标准能力,遇到的跨作用域报错是tf.function装饰器使用方式不符合执行逻辑导致的,不是add_loss本身的功能限制。
问题复现信息
- 测试环境:Python 3.9.12,TensorFlow 2.8.0
- 复现代码:
import tensorflow as tf from tensorflow import keras def main(): layer = ActivityRegularizationLayer() inputs = tf.constant(5.) with tf.GradientTape() as tape: y = layer(inputs) loss = tf.reduce_mean(y) loss += sum(layer.losses) grad = tape.gradient(loss, layer.trainable_weights) print(f"loss={float(loss)}, grad={grad}") class ActivityRegularizationLayer(keras.layers.Layer): def __init__(self, rate=1e-2): super().__init__() self.rate = rate @tf.function def call(self, inputs): self.add_loss(self.rate * tf.reduce_sum(inputs)) return inputs
- 报错信息:
The tensor <tf.Tensor 'mul:0' shape=() dtype=float32> cannot be accessed from here, because it was defined in FuncGraph(name=call, id=46917885252656), which is out of scope.
- 对照测试结果:
- 移除
call方法的@tf.function装饰器,代码正常运行,输出:loss=5.050000190734863, grad=[] - 删除总损失计算中
sum(layer.losses)的代码行,代码正常运行,输出:loss=5.0, grad=[]
- 移除
报错原因
给call方法单独添加@tf.function装饰器后,方法内部的所有运算会被封装到独立的FuncGraph静态图中执行,add_loss注册的损失张量属于这个内部静态图的局部资源,不会自动暴露到外层的Eager执行上下文。
此时在外层Eager环境下直接读取layer.losses获取损失张量,会触发跨FuncGraph的资源访问,外层上下文没有权限访问内部图的节点,就会抛出超出作用域的错误。
两种不报错的场景逻辑也很明确:
- 移除
@tf.function时,call默认在Eager模式下执行,所有张量都处于同一个全局执行上下文,不存在跨作用域问题 - 不读取
layer.losses时,不会触发跨图资源访问,自然不会报错
解决方法
可以根据使用场景选择以下任意一种方案:
- 不要单独给
call方法加@tf.function装饰器。如果需要图执行加速,把包含前向传播、损失计算、梯度求导的整个训练逻辑包裹在同一个@tf.function下,保证所有运算处于同一个FuncGraph作用域内,示例代码如下:
import tensorflow as tf from tensorflow import keras class ActivityRegularizationLayer(keras.layers.Layer): def __init__(self, rate=1e-2): super().__init__() self.rate = rate def call(self, inputs): self.add_loss(self.rate * tf.reduce_sum(inputs)) return inputs @tf.function def train_step(layer, inputs): with tf.GradientTape() as tape: y = layer(inputs) loss = tf.reduce_mean(y) loss += sum(layer.losses) grad = tape.gradient(loss, layer.trainable_weights) return loss, grad # 执行 layer = ActivityRegularizationLayer() inputs = tf.constant(5.) loss, grad = train_step(layer, inputs) print(f"loss={float(loss)}, grad={grad}")
- 如果一定要给
call方法加@tf.function,就不要在外层Eager上下文手动读取layer.losses。优先使用Keras内置的fit训练流程,或者在自定义Model的train_step方法内完成损失收集计算,框架会自动在匹配的作用域内获取add_loss注册的损失,不会触发跨域问题。 - 升级TensorFlow到2.10及以上版本,高版本对
tf.function下的add_loss张量跨上下文访问做了兼容处理,多数场景下不会再抛出该类作用域错误。
内容的提问来源于stack exchange,提问作者LexTron
相关产品推荐
相关产品推荐

