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

自定义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

说明

  1. 移除了normalizer相关的实例属性和初始化参数,因为你的场景下每个样本的归一化项就是自身的y_true,不需要复用或保存这个值。
  2. 简化了get_config方法,避免尝试访问图内张量。
  3. 保留了原有的张量兼容性处理、维度调整逻辑,确保指标能正确处理不同形状的输入。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 20:32:39