自定义ResidualPerUnitObservation指标报错:张量超出作用域求助
解决自定义Keras指标跨作用域张量访问错误
错误原因
你遇到的问题根源有两点:
- 在
update_state方法中,将属于训练函数图内的y_true张量赋值给了指标类的实例属性self.normalizer,Keras会尝试序列化指标的实例属性,导致跨FuncGraph访问张量的错误。 get_config方法中尝试评估self.normalizer张量,图模式下不允许在图外部访问图内张量。
修正后的代码
既然你的需求是计算abs(y_true - y_pred)/y_true,不需要将y_true保存为实例属性,直接在每次更新状态时使用当前的y_true作为归一化项即可:
class ResidualPerUnitObservation(base_metric.Mean): @dtensor_utils.inject_mesh def __init__(self, name=None, dtype=None): # 移除normalizer参数,不需要保存实例属性 super().__init__(name=name, dtype=dtype) def update_state(self, y_true, y_pred, sample_weight=None): y_true = tf.cast(y_true, self._dtype) y_pred = tf.cast(y_pred, self._dtype) [y_pred, y_true], sample_weight = metrics_utils.ragged_assert_compatible_and_get_flat_values( [y_pred, y_true], sample_weight ) y_pred, y_true = losses_utils.squeeze_or_expand_dimensions(y_pred, y_true) y_pred.shape.assert_is_compatible_with(y_true.shape) # 直接用当前的y_true作为归一化项,不需要存为实例属性 relative_errors = tf.math.divide_no_nan(tf.abs(y_true - y_pred), y_true) return super().update_state(relative_errors, sample_weight=sample_weight) def get_config(self): # 不需要处理normalizer,直接返回父类配置 base_config = super().get_config() return base_config
说明
- 移除了
normalizer相关的实例属性和初始化参数,因为你的场景下每个样本的归一化项就是自身的y_true,不需要复用或保存这个值。 - 简化了
get_config方法,避免尝试访问图内张量。 - 保留了原有的张量兼容性处理、维度调整逻辑,确保指标能正确处理不同形状的输入。
内容的提问来源于stack exchange,提问作者Jacky
相关产品推荐
相关产品推荐

