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

TensorFlow Federated自定义准确率报AttributeError解决方法

报错根因

tff.learning.from_keras_model对传入的metrics参数有明确要求:仅接受继承自tf.keras.metrics.Metric的指标实例,不支持直接传入普通Python函数。
当前传入的hinge_accuracy是无状态的普通函数,没有Keras指标标准实现中用来存储跨批次统计值的variables属性。TFF在执行本地指标输出逻辑时,会遍历所有指标的variables属性读取累积的统计结果做联邦聚合,访问函数对象的该属性时就会抛出'function' object has no attribute 'variables'异常。

排查步骤
  • 校验传入指标的类型:执行print(type(hinge_accuracy))会输出<class 'function'>,不符合TFF要求的Metric子类实例规范。
  • 核对接口逻辑:和本地Kerasmodel.compile可兼容裸指标函数的逻辑不同,联邦学习场景下指标需要持有可跨批次、跨通信轮持久化的状态变量,才能支持多客户端的指标聚合,无状态函数无法满足该要求。
修复方案

将裸函数形式的自定义指标改写为继承tf.keras.metrics.Metric的有状态指标类,实现状态初始化、批次更新、结果计算、状态重置四个核心逻辑即可,修改后的可运行代码如下:

import tensorflow as tf
import tensorflow_federated as tff

class HingeAccuracy(tf.keras.metrics.Metric):
    def __init__(self, name='hinge_accuracy', **kwargs):
        super().__init__(name=name, **kwargs)
        # 定义两个状态变量,分别累计预测正确的样本数、总样本数
        self.correct_count = self.add_weight(name='correct_cnt', initializer='zeros')
        self.total_count = self.add_weight(name='total_cnt', initializer='zeros')

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 原函数的判断逻辑迁移到状态更新步骤
        y_true_label = tf.squeeze(y_true) > 0.0
        y_pred_label = tf.squeeze(y_pred) > 0.0
        batch_correct = tf.reduce_sum(tf.cast(y_true_label == y_pred_label, tf.float32))
        batch_total = tf.cast(tf.size(y_true_label), tf.float32)

        # 累加更新状态
        self.correct_count.assign_add(batch_correct)
        self.total_count.assign_add(batch_total)

    def result(self):
        # 计算并返回最终指标值
        return self.correct_count / self.total_count

    def reset_state(self):
        # 每轮本地评估/训练开始前重置状态
        self.correct_count.assign(0.0)
        self.total_count.assign(0.0)

def model_fn():
    keras_model_clone = create_baseline_model()
    return tff.learning.from_keras_model(
        keras_model_clone,
        input_spec=preprocessed_example_dataset.element_spec,
        loss=tf.keras.losses.Hinge(),
        # 传入自定义指标的实例,而非裸函数
        metrics=[HingeAccuracy()]
    )

快捷方案:如果不需要自定义实现,直接使用Keras内置的tf.keras.metrics.BinaryAccuracy(threshold=0.0)可以实现和原有hinge_accuracy完全一致的计算逻辑,无需额外编写自定义类,直接传入该实例即可。

修改完成后重新运行,TFF即可正常读取指标实例的变量完成本地统计和联邦聚合,该AttributeError不会复现。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 12:48:15